diff --git a/.gitignore b/.gitignore index 5b23f68f..2894211c 100644 --- a/.gitignore +++ b/.gitignore @@ -23,6 +23,7 @@ bin/test ### IntelliJ IDEA ### .idea +.junie *.iws *.iml *.ipr diff --git a/config/checkstyle/intellij_codestyle.xml b/config/checkstyle/intellij_codestyle.xml new file mode 100644 index 00000000..77cc7369 --- /dev/null +++ b/config/checkstyle/intellij_codestyle.xml @@ -0,0 +1,52 @@ + + + + + + + + + + + + \ No newline at end of file diff --git a/solvers/circuitprocessing/circuitexecution/executor.py b/solvers/circuitprocessing/circuitexecution/executor.py new file mode 100644 index 00000000..379288f1 --- /dev/null +++ b/solvers/circuitprocessing/circuitexecution/executor.py @@ -0,0 +1,40 @@ +import sys +from pytket.qasm import circuit_from_qasm_str + +input_path = sys.argv[1] +num_runs = int(sys.argv[2]) +backend_name = sys.argv[3] + +with open(input_path, 'r') as input_file: + text = input_file.read() + +try: + circuit = circuit_from_qasm_str(text) +except Exception as e: + print("Was not able to convert to OpenQASM: ", e) + sys.exit(1) + +if backend_name == "aer": + from pytket.extensions.qiskit import AerBackend + backend = AerBackend() +elif backend_name == "qulacs": + from pytket.extensions.qulacs import QulacsBackend + backend = QulacsBackend() +elif backend_name == "aer_noisy": + from pytket.extensions.qiskit import AerBackend + from qiskit_aer.noise import NoiseModel + from qiskit_aer.noise.errors import depolarizing_error + # https://docs.quantinuum.com/tket/user-guide/manual/manual_noise.html + noise_model = NoiseModel() + noise_model.add_readout_error([[0.9, 0.1], [0.1, 0.9]], [0]) + noise_model.add_readout_error([[0.95, 0.05], [0.05, 0.95]], [1]) + noise_model.add_quantum_error(depolarizing_error(0.1, 2), ["cx"], [0, 1]) + backend = AerBackend(noise_model) +else: + print(f"Unknown backend: {backend_name}", file=sys.stderr) + sys.exit(1) + +c = backend.get_compiled_circuit(circuit) +handle = backend.process_circuit(c, n_shots=num_runs) +counts = backend.get_result(handle).get_counts() +print(counts) diff --git a/solvers/circuitprocessing/circuitexecution/requirements.txt b/solvers/circuitprocessing/circuitexecution/requirements.txt new file mode 100644 index 00000000..46f05119 --- /dev/null +++ b/solvers/circuitprocessing/circuitexecution/requirements.txt @@ -0,0 +1,3 @@ +pytket +pytket-qiskit +# pytket-qulacs \ No newline at end of file diff --git a/solvers/circuitprocessing/circuitoptimizing/decompose-multi-cx/decompose_multi_cx_optimizer.py b/solvers/circuitprocessing/circuitoptimizing/decompose-multi-cx/decompose_multi_cx_optimizer.py new file mode 100644 index 00000000..3378a019 --- /dev/null +++ b/solvers/circuitprocessing/circuitoptimizing/decompose-multi-cx/decompose_multi_cx_optimizer.py @@ -0,0 +1,25 @@ +import sys +from pytket.qasm import circuit_from_qasm_str, circuit_to_qasm_str +from pytket.predicates import CompilationUnit +from pytket.passes import DecomposeMultiQubitsCX + +input_path = sys.argv[1] + +# read input from file +with open(input_path, 'r') as input_file: + text = input_file.read() + +input_circuit = text + +try: + circuit = circuit_from_qasm_str(input_circuit) +except Exception as e: + print("Was not able to convert to OpenQASM: ", e) + sys.exit(1) + +pass1 = DecomposeMultiQubitsCX() +cu = CompilationUnit(circuit) +pass1.apply(cu) + +print(circuit_to_qasm_str(cu.circuit)) + diff --git a/solvers/circuitprocessing/circuitoptimizing/remove-redundancies/remove_redundancies_optimizer.py b/solvers/circuitprocessing/circuitoptimizing/remove-redundancies/remove_redundancies_optimizer.py new file mode 100644 index 00000000..a238a84a --- /dev/null +++ b/solvers/circuitprocessing/circuitoptimizing/remove-redundancies/remove_redundancies_optimizer.py @@ -0,0 +1,24 @@ +import sys +from pytket.qasm import circuit_from_qasm_str, circuit_to_qasm_str +from pytket.predicates import CompilationUnit +from pytket.passes import RemoveRedundancies + +input_path = sys.argv[1] + +# read input from file +with open(input_path, 'r') as input_file: + text = input_file.read() + +input_circuit = text + +try: + circuit = circuit_from_qasm_str(input_circuit) +except Exception as e: + print("Was not able to convert to OpenQASM: ", e) + sys.exit(1) + +pass1 = RemoveRedundancies() +cu = CompilationUnit(circuit) +pass1.apply(cu) + +print(circuit_to_qasm_str(cu.circuit)) \ No newline at end of file diff --git a/solvers/circuitprocessing/circuitoptimizing/requirements.txt b/solvers/circuitprocessing/circuitoptimizing/requirements.txt new file mode 100644 index 00000000..f32bd3e4 --- /dev/null +++ b/solvers/circuitprocessing/circuitoptimizing/requirements.txt @@ -0,0 +1 @@ +pytket \ No newline at end of file diff --git a/solvers/qiskit/knapsack_quantum_tree_generator/BaseCircuit/adder.py b/solvers/qiskit/knapsack_quantum_tree_generator/BaseCircuit/adder.py new file mode 100644 index 00000000..e87bf1ce --- /dev/null +++ b/solvers/qiskit/knapsack_quantum_tree_generator/BaseCircuit/adder.py @@ -0,0 +1,32 @@ +from qiskit import QuantumCircuit +from qiskit.circuit import Gate +from qiskit.circuit.library import QFT +from math import pi +from typing import Optional + +def constant_qft_add_gate(n_bits: int, const: int, name: Optional[str] = None) -> Gate: + """ + QFT-based constant adder: |v> -> |(v + const) mod 2^n_bits>. + Uses QFT, per-qubit phase rotations, and inverse QFT. + """ + + qc = QuantumCircuit(n_bits, name=name or f"AddConst({const})") + qft_gate = QFT(n_bits, do_swaps=False).to_gate(label="QFT") + iqft_gate = qft_gate.inverse() + + qc.append(qft_gate, range(n_bits)) + for k in range(n_bits): + angle = 2 * pi * const / (2 ** (k + 1)) + qc.p(angle, k) + qc.append(iqft_gate, range(n_bits)) + return qc.to_gate(label=name or f"AddConst({const})") + + + +def constant_qft_sub_gate(n_bits: int, const: int, name: Optional[str] = None) -> Gate: + """ + QFT-based constant subtractor: |v> -> |(v - const) mod 2^n_bits>. + Implemented by flipping the signs of the adder’s phase rotations [1]. + """ + return constant_qft_add_gate(n_bits, const=-const, name=name or f"SubConst({const})") + diff --git a/solvers/qiskit/knapsack_quantum_tree_generator/BaseCircuit/circuit_core.py b/solvers/qiskit/knapsack_quantum_tree_generator/BaseCircuit/circuit_core.py new file mode 100644 index 00000000..e3f59b08 --- /dev/null +++ b/solvers/qiskit/knapsack_quantum_tree_generator/BaseCircuit/circuit_core.py @@ -0,0 +1,195 @@ +from typing import Dict, List, Sequence, Any, Union, Optional +from qiskit import QuantumCircuit, QuantumRegister, AncillaRegister +from qiskit.circuit import Instruction, ClassicalRegister +from knapsack.knapsack import KnapsackInstance + + + +class RegisterBank: + """ + Helper to keep named registers organized. + """ + + def __init__(self): + self.q: Dict[str, QuantumRegister] = {} + self.c: Dict[str, ClassicalRegister] = {} + self.a: Dict[str, AncillaRegister] = {} + + def add_qubits(self, name: str, size: int) -> QuantumRegister: + if name in self.q: + raise ValueError(f"Quantum register '{name}' already exists") + reg = QuantumRegister(size, name=name) + self.q[name] = reg + return reg + + def add_ancilla(self, name: str, size: int) -> AncillaRegister: + if name in self.a: + raise ValueError(f"Ancilla register '{name}' already exists") + reg = AncillaRegister(size, name=name) + self.a[name] = reg + return reg + + def add_clbits(self, name: str, size: int) -> ClassicalRegister: + if name in self.c: + raise ValueError(f"Classical register '{name}' already exists") + reg = ClassicalRegister(size, name=name) + self.c[name] = reg + return reg + + def get(self, name: List[str]) -> Union[QuantumRegister, AncillaRegister, ClassicalRegister]: + list_of_registers = [] + for n in name: + if n in self.q: + list_of_registers.append(self.q[n]) + if n in self.a: + list_of_registers.append(self.a[n]) + if n in self.c: + list_of_registers.append(self.c[n]) + if len(list_of_registers) == 1: + return list_of_registers[0] + elif len(list_of_registers) > 1: + return list_of_registers + else: + raise KeyError(f"Register '{name}' not found") + + def has_measurements(self) -> bool: + return len(self.c) > 0 + + +class Circuit: + """ + Thin wrapper around QuantumCircuit with named register management. + """ + + def __init__(self, name: str = "circuit"): + self.name = name + self.registers = RegisterBank() + self.qc = QuantumCircuit(name=name) + self.metadata: Dict[str, Any] = {} + + def add_qubits(self, name: str, size: int) -> QuantumRegister: + reg = self.registers.add_qubits(name, size) + self.qc.add_register(reg) + return reg + + def add_ancilla(self, name: str, size: int) -> AncillaRegister: + reg = self.registers.add_ancilla(name, size) + self.qc.add_register(reg) + return reg + + def add_clbits(self, name: str, size: int) -> ClassicalRegister: + reg = self.registers.add_clbits(name, size) + self.qc.add_register(reg) + return reg + + def get_register(self, name: Union[str, List[str]]) -> Union[QuantumRegister, AncillaRegister, ClassicalRegister]: + """Get a register by name.""" + if isinstance(name, str): + return self.registers.get([name]) + return self.registers.get(name) + + def get_qubits_in_registers(self, names: List[str] = [], all=False) -> List[Any]: + """Get all qubits in the specified registers as a flat list.""" + if all and len(names) != 0: + raise ValueError("If 'all' is True, 'names' must be an empty list.") + if not all and len(names) == 0: + raise ValueError("If 'all' is False, 'names' must contain at least one register name.") + + if all: + names = list(self.registers.q.keys()) + list(self.registers.a.keys()) + list(self.registers.c.keys()) + qubits = [] + for name in names: + reg = self.get_register(name) + for q in reg: + qubits.append(q) + return qubits + + + def append_instruction( + self, + inst: Instruction, + qargs: Sequence[Any], + cargs: Optional[Sequence[Any]] = None, + ): + self.qc.append(inst, qargs, [] if cargs is None else cargs) + + def append_subcircuit_as_instruction( + self, + sub: QuantumCircuit, + qubits: Sequence[Any], + clbits: Optional[Sequence[Any]] = None, + name: Optional[str] = None, + ): + """Append a subcircuit as a single instruction to this circuit.""" + inst = sub.to_instruction() + if name: + inst.name = name + self.append_instruction(inst, qubits, [] if clbits is None else clbits) + + + + def append_subcircuit_inline_simple(self, sub: QuantumCircuit): + """ + Adding instruction 1 by 1 without mapping, assuming: + - both circuits have the same number of qubits, + - no classical bits are present. + """ + if len(sub.qc.clbits) != 0 or len(self.qc.clbits) != 0: + raise ValueError("This simplified method requires circuits without classical bits.") + if len(sub.qc.qubits) != len(self.qc.qubits): + raise ValueError("Both circuits must have the same number of qubits.") + + # Map sub qubits -> target qubits by index + qmap = {sub_q: self.qc.qubits[i] for i, sub_q in enumerate(sub.qc.qubits)} + + # Append each instruction as-is + for ci in sub.qc.data: + op = ci.operation + mapped_qargs = [qmap[q] for q in ci.qubits] + self.qc.append(op, mapped_qargs, []) + + def to_instruction(self, name: Optional[str] = None) -> Instruction: + inst = self.qc.to_instruction() + if name: + inst.name = name + return inst + + def prepare_knapsack_circuit(self, knapsack_instance: 'KnapsackInstance', has_ancillas = False) -> None: + """Prepare the quantum circuit for the knapsack problem.""" + self.add_qubits("items", knapsack_instance.num_items) + self.add_qubits("capacity", knapsack_instance.capacity.bit_length()) + max_profit = sum(it.value for it in knapsack_instance.items) + self.add_qubits("profit", max_profit.bit_length()) + self.add_ancilla("oracle", 1) + if has_ancillas: + self.add_ancilla("compare_flag", 1) + self.add_ancilla("comparator", max(knapsack_instance.num_items, knapsack_instance.capacity.bit_length(), max_profit.bit_length())) + # set capacity register to knapsack capacity + + + def measure_items(self, creg_name: str = "items_c"): + self.add_clbits(creg_name, len(self.get_register("items"))) + # Freeze ordering before measuring + self.qc.barrier(*self.qc.qubits) + self.measure_register("items", creg_name) + + def measure_register(self, qreg_name: str, creg_name: str): + qreg = self.get_register(qreg_name) + creg = self.get_register(creg_name) + self.qc.measure(qreg, creg) + + + def draw(self, output: str = "text") -> Any: + return self.qc.draw(output=output) + + def get_circuit_as_instruction(self, name: Optional[str] = None) -> Instruction: + inst = self.qc.to_instruction() + if name: + inst.name = name + return inst + + + + + + diff --git a/solvers/qiskit/knapsack_quantum_tree_generator/BaseCircuit/comparator.py b/solvers/qiskit/knapsack_quantum_tree_generator/BaseCircuit/comparator.py new file mode 100644 index 00000000..15672e37 --- /dev/null +++ b/solvers/qiskit/knapsack_quantum_tree_generator/BaseCircuit/comparator.py @@ -0,0 +1,109 @@ +from qiskit import QuantumCircuit +from qiskit.circuit import Instruction +from typing import Sequence, Optional, Union +from qiskit.circuit import QuantumRegister, Qubit + + +def apply_threshold_controlled_U(circ: QuantumCircuit, + qreg: Sequence, + target, + U: Instruction, + w: int, + inplace: bool = True, + broadcast: bool = False, + label: Optional[str] = None) -> Union[QuantumCircuit, Instruction]: + """ + There are 3 cases: qreg = w, qreg < w, qreg > w. + Choose the first position where qreg and w differ when scanning from MSB to LSB --> position j. + - (1) If qreg[j] = 0 and w[j] = 1 --> qreg < w --> do nothing. + - (2) If qreg[j] = 1 and w[j] = 0 --> qreg > w --> apply U on target. + - (3) if no such j is found --> qreg = w --> apply U on target. + + We only care about cases (2) and (3). So for all j where w[j] = 0, we control on all qreg[k>j] = w[k] and qreg[j] = 1 to apply U on target for case (2). + Finally, we control on all qreg[k] = w[k] to apply U on target for case (3). + Because if qreg[k>j] = w[k] not true there must be some index t, with qreg[t]=0 und w[t]=1, which is case (1) and we don't care about it. + + Since we at most check all bits, we have O(log(C)) gates. + + If inplace=True: + - Modifies 'circ' in place and returns 'circ'. + If inplace=False: + - Returns an Instruction (built from a fresh (n+1)-qubit subcircuit) that performs + the same operation, with name/label 'label' or 'CU(c>={w})'. + You must append it with the qubits: list(qreg) + [target]. + + Broadcast applies U to all targets in target register + """ + n = len(qreg) + if U.num_qubits != 1: + raise ValueError("U must be a single-qubit gate/instruction.") + + if broadcast: + if not isinstance(target, QuantumRegister): + raise ValueError("When broadcast=True, 'target' must be a non-empty sequence of qubits.") + tgt_list = list(target) + else: + # Single target + if isinstance(target, Sequence): + raise ValueError("When broadcast=False, 'target' must be a single qubit (or a sequence of length 1).") + else: + tgt_list = [target] + + m = len(tgt_list) + + def _body(on_circ: QuantumCircuit, reg, targets): + # Edge cases on the comparator condition + if n == 0: + if w <= 0: + for t in targets: + on_circ.append(U, [t]) + return + if w <= 0: + for t in targets: + on_circ.append(U, [t]) + return + if w >= (1 << n): + return + + # Bits of w in little-endian: bits[i] is bit i (i=0 is LSB, i=n-1 is MSB) + bits = [(w >> i) & 1 for i in range(n)] + + # Clauses for first-difference positions where w_j = 0 (x >= w because x_j=1 and all higher bits equal) + for j in range(n - 1, -1, -1): # MSB down to LSB + if bits[j] == 0: + zeros_higher = [t for t in range(j + 1, n) if bits[t] == 0] + for t in zeros_higher: + on_circ.x(reg[t]) # control-on-0 -> control-on-1 for higher positions + control_qubits = [reg[t] for t in range(j + 1, n)] + [reg[j]] + CU = U.control(len(control_qubits)) + for tgt in targets: + on_circ.append(CU, control_qubits + [tgt]) + for t in zeros_higher: + on_circ.x(reg[t]) # undo + + # Equality clause x == w (all bits match w) + zeros_all = [t for t in range(n) if bits[t] == 0] + for t in zeros_all: + on_circ.x(reg[t]) + CU_eq = U.control(n) + for tgt in targets: + on_circ.append(CU_eq, list(reg) + [tgt]) + for t in zeros_all: + on_circ.x(reg[t]) + + if inplace: + _body(circ, qreg, tgt_list) + return circ + else: + # Build a minimal subcircuit and return as instruction + if broadcast: + sc = QuantumCircuit(n + m, name=label or f"CU_broadcast(c>={w})") + reg = [sc.qubits[i] for i in range(n)] + targets = [sc.qubits[n + i] for i in range(m)] + else: + sc = QuantumCircuit(n + 1, name=label or f"CU(c>={w})") + reg = [sc.qubits[i] for i in range(n)] + targets = [sc.qubits[n]] # single target + _body(sc, reg, targets) + return sc.to_instruction() + diff --git a/solvers/qiskit/knapsack_quantum_tree_generator/BaseCircuit/simulator.py b/solvers/qiskit/knapsack_quantum_tree_generator/BaseCircuit/simulator.py new file mode 100644 index 00000000..f11d17ed --- /dev/null +++ b/solvers/qiskit/knapsack_quantum_tree_generator/BaseCircuit/simulator.py @@ -0,0 +1,103 @@ + + +from typing import Any, Dict, Optional +from data.SimulateOptions import SimulationOptions +from BaseCircuit.circuit_core import Circuit +from qiskit import transpile +from qiskit.quantum_info import Statevector +from qiskit.visualization import plot_histogram + +try: + # Qiskit Aer is optional; for shot-based or noisy simulation + from qiskit_aer import AerSimulator + AER_AVAILABLE = True +except Exception: + AER_AVAILABLE = False + + +class Simulator: + def __init__(self, circuit: Circuit): + self.circuit = circuit + self.result = None + self.simulated = False + + def simulate(self, sim: Optional[SimulationOptions] = None) -> Dict[str, Any]: + """Start simulation of the circuit according to SimulationOptions.""" + sim = sim or SimulationOptions() + + # Analytic statevector path (no noise, no shots) + if sim.method == "statevector" and sim.noise_model is None and (sim.shots is None or sim.shots <= 0): + sv = Statevector.from_instruction(self.qc) + return {"statevector": sv} + + # Aer required for shot-based / noisy runs + if not AER_AVAILABLE: + raise RuntimeError("Aer is not available. Install qiskit-aer for shot-based/noisy simulation.") + + # Default shots when not in analytic statevector mode + if sim.shots is None: + sim.shots = 1024 + + qc_to_run = self.circuit.qc.copy() + + # Auto-insert measurements only if requested + if sim.method != "statevector" and sim.shots and sim.shots > 0 and sim.measure_all_if_none: + def _has_meas_or_cond(qc): + for ci in qc.data: + op = ci.operation # CircuitInstruction.operation + if op.name == "measure": + return True + if getattr(ci, "condition", None) is not None: + return True + return False + + if not _has_meas_or_cond(qc_to_run): + qc_to_run.measure_all() + + method = sim.method if sim.method in ("automatic", "statevector", "density_matrix") else "automatic" + backend = AerSimulator(method=method) + if sim.noise_model is not None: + backend.set_options(noise_model=sim.noise_model) + if sim.seed_simulator is not None: + backend.set_options(seed_simulator=sim.seed_simulator) + + tqc = transpile(qc_to_run, backend, optimization_level=0) + job = backend.run(tqc, shots=sim.shots) + result = job.result() + + payload: Dict[str, Any] = {"result": result} + + # Counts, invert bitstrings to upset qiskit MSB notation. + try: + counts = result.get_counts(0) + counts = {bits[::-1]: cnt for bits, cnt in counts.items()} + payload["counts"] = counts + except Exception: + pass + + # Statevector + try: + payload["statevector"] = result.get_statevector(0) + except Exception: + pass + + self.result = payload + self.simulated = True + return True + + def get_probabilities_from_counts(self): + if self.simulated is False or self.result is None or "counts" not in self.result: + raise RuntimeError("No counts available. Please run simulate() with shot-based options first.") + total = sum(self.result["counts"].values()) + probs = {bits: c / total for bits, c in self.result["counts"].items()} + # Sort by probability (descending) and print + return {bits: p for bits, p in sorted(probs.items(), key=lambda kv: kv[1], reverse=True)} + + def plot_counts(self) -> Any: + """Plot counts using Qiskit visualization tools.""" + return plot_histogram(data=self.result["counts"], title="Measurement Counts") + + def plot_infered_prob_from_counts(self) -> Any: + """Plot inferred probabilities from counts using Qiskit visualization tools.""" + return plot_histogram(data=self.get_probabilities_from_counts(), title="Inferred Probabilities from Counts", sort='value_desc') + diff --git a/solvers/qiskit/knapsack_quantum_tree_generator/BaseCircuit/statevector_debug_util.py b/solvers/qiskit/knapsack_quantum_tree_generator/BaseCircuit/statevector_debug_util.py new file mode 100644 index 00000000..97c264e8 --- /dev/null +++ b/solvers/qiskit/knapsack_quantum_tree_generator/BaseCircuit/statevector_debug_util.py @@ -0,0 +1,74 @@ +from pathlib import Path +import numpy as np +from qiskit import QuantumCircuit, transpile +from qiskit_aer import AerSimulator +import json +from qiskit.quantum_info import Statevector + + +def save_step_by_step_sv_probs( + qc: QuantumCircuit, + out_dir: str = "log/full_probs_by_step", + threshold: float = 1e-6, +): + """ + Save the full probability distribution after every gate to separate JSON files. + Each JSON maps bitstring -> probability for the entire system. + """ + qc_dbg = QuantumCircuit(*qc.qregs, *qc.cregs) + for i, (op, qargs, cargs) in enumerate(qc.data): + qc_dbg.append(op, qargs, cargs) + qc_dbg.save_statevector(label=f"sv_{i:04d}") + + backend = AerSimulator(method="statevector") + tqc = transpile(qc_dbg, backend, optimization_level=0) + result = backend.run(tqc).result() + + out_path = Path(out_dir) + out_path.mkdir(parents=True, exist_ok=True) + + data = result.data(0) + for key in sorted(k for k in data.keys() if k.startswith("sv_")): + sv = data[key] + if not isinstance(sv, Statevector): + sv = Statevector(sv) + probs = sv.probabilities_dict() # all qubits + # Filter small entries for readability + probs = {k: float(v) for k, v in probs.items() if v > threshold} + with open(out_path / f"{key}.json", "w") as f: + json.dump(probs, f, indent=2) + print(f"Wrote full probability snapshots to: {out_dir}") + +def save_step_by_step_sv_amplitudes( + qc: QuantumCircuit, + out_dir: str = "log/full_probs_by_step", + threshold: float = 1e-6, +): + """ + Save the full amplitude distribution after every gate to separate JSON files. + Each JSON maps bitstring -> {re, im}. + """ + qc_dbg = QuantumCircuit(*qc.qregs, *qc.cregs) + for i, (op, qargs, cargs) in enumerate(qc.data): + qc_dbg.append(op, qargs, cargs) + qc_dbg.save_statevector(label=f"sv_{i:04d}") + backend = AerSimulator(method="statevector") + tqc = transpile(qc_dbg, backend, optimization_level=0) + result = backend.run(tqc).result() + out_path = Path(out_dir) + out_path.mkdir(parents=True, exist_ok=True) + data = result.data(0) + for key in sorted(k for k in data.keys() if k.startswith("sv_")): + sv = data[key] + if not isinstance(sv, Statevector): + sv = Statevector(sv) + amp_dict = sv.to_dict() # bitstring -> complex amplitude + # Filter small entries for readability (by probability magnitude) + amps = { + k: {"re": float(np.real(v)), "im": float(np.imag(v))} + for k, v in amp_dict.items() + if (np.abs(v) ** 2) > threshold + } + with open(out_path / f"{key}.json", "w") as f: + json.dump(amps, f, indent=2) + print(f"Wrote full amplitude snapshots to: {out_dir}") diff --git a/solvers/qiskit/knapsack_quantum_tree_generator/Grover/grover.py b/solvers/qiskit/knapsack_quantum_tree_generator/Grover/grover.py new file mode 100644 index 00000000..1993e758 --- /dev/null +++ b/solvers/qiskit/knapsack_quantum_tree_generator/Grover/grover.py @@ -0,0 +1,158 @@ +from qiskit import QuantumCircuit +from BaseCircuit.circuit_core import Circuit +from knapsack.knapsack import KnapsackInstance +from typing import Tuple +from BaseCircuit.comparator import apply_threshold_controlled_U +from qiskit.circuit.library import XGate + + +class Grover(Circuit): + def __init__( + self, + input_circuit: Circuit, + knapsack: KnapsackInstance, + depth_interval: Tuple[int, int] = (0, -1), + name: str = "Grover", + current_threshold: int = 0, + optimal_iterations: int = None + ): + """ + Initialize the Grover instance. + """ + self.input_circuit = input_circuit + self.name = name + self.knapsack = knapsack + self.current_circuit = self.copy_circuit_structure() + self.current_threshold = current_threshold + #the input circuit is always also the state prep circuit + self.state_prep_instruction = self.input_circuit.get_circuit_as_instruction(name="QTG" if depth_interval==(0,-1) else "QTG_partial") + if optimal_iterations is None: + raise ValueError("optimal_iterations must be provided for Grover initialization.") + self.optimal_iterations = optimal_iterations + + + self.depth_interval = depth_interval + if depth_interval[1] == -1: + self.depth_interval = (depth_interval[0], len(knapsack.items)) + + # this modifies the current threshold to exlude items than canont be feasible from some partial depth + if self.depth_interval[1] < len(knapsack.items): + self.current_threshold = self.current_threshold - sum(self.knapsack.values[self.depth_interval[1]:]) + print(f'Grover inner threshold modified to {self.current_threshold} for depth interval {self.depth_interval}') + + + + def copy_circuit_structure(self, name:str = None) -> QuantumCircuit: + qc = Circuit(name=name if name is not None else self.name) + for qreg in self.input_circuit.registers.q: + reg = self.input_circuit.get_register(qreg) + qc.add_qubits(reg.name, reg.size) + for areg in self.input_circuit.registers.a: + reg = self.input_circuit.get_register(areg) + qc.add_ancilla(reg.name, reg.size) + for creg in self.input_circuit.registers.c: + reg = self.input_circuit.get_register(creg) + qc.add_clbits(reg.name, reg.size) + return qc + + def oracle(self, + profit_reg_name: str = "profit", + oracle_reg_name: str = "oracle", + as_instruction: bool = False, + label: str | None = None): + + qc = self.current_circuit.qc + profit_reg = self.current_circuit.get_register(profit_reg_name) + oracle = self.current_circuit.get_register(oracle_reg_name)[0] + + if not as_instruction: + qc.h(oracle) + qc.z(oracle) + + # apply phase to ancilla which will become a global phase for the decision qubits + controlled_Z = apply_threshold_controlled_U( + circ=qc, + qreg=profit_reg, + target=oracle, + U=XGate(), + w=self.current_threshold, + broadcast=False, + inplace=False + ) + self.current_circuit.append_instruction(controlled_Z, list(profit_reg)+ [oracle]) + # Unprepare ancilla back to |0> + qc.z(oracle) + qc.h(oracle) + return qc + else: + # Build a reusable instruction that includes prepare-|->, comparator, and unprepare + n = len(profit_reg) + sub = QuantumCircuit(n + 1, name=label if label is not None else f"Oracle_P(x)>={self.current_threshold}") + reg = [sub.qubits[i] for i in range(n)] + anc = sub.qubits[n] + + sub.h(anc) + sub.z(anc) + apply_threshold_controlled_U( + circ=sub, + qreg=reg, + target=anc, + U=XGate(), + w=self.current_threshold, + broadcast=False, + inplace=True + ) + sub.z(anc) + sub.h(anc) + + instr = sub.to_instruction() + qc.append(instr, list(profit_reg) + [oracle]) + return instr + + def diffuser(self, as_instruction: bool = False, label: str | None = None): + """ + Build/apply the Grover reflection R_psi = A (I - 2|0...0><0...0|) A^\dagger, + where A is given by self.state_prep_circuit. + + If as_instruction=True, rappend whole instruction. + If as_instruction=False, append to the same qubits that state_prep_circuit prepares. + """ + qc = self.current_circuit.qc + if not as_instruction: + # the x gates transfer to controls on 1, the hadamard mcx hadarmard is a mcz so we control the last one phase flip on 11111111, and Z only acts on 1. + qc.append(self.state_prep_instruction.inverse(), list(qc.qubits)[:self.state_prep_instruction.num_qubits]) + for q in list(qc.qubits)[:self.state_prep_instruction.num_qubits-1]: + qc.x(q) + qc.x(list(qc.qubits)[self.state_prep_instruction.num_qubits - 1]) + qc.h(list(qc.qubits)[self.state_prep_instruction.num_qubits - 1]) + qc.mcx(list(qc.qubits)[:self.state_prep_instruction.num_qubits - 1], qc.qubits[self.state_prep_instruction.num_qubits - 1]) + qc.h(list(qc.qubits)[self.state_prep_instruction.num_qubits - 1]) + qc.x(list(qc.qubits)[self.state_prep_instruction.num_qubits - 1]) + for q in list(qc.qubits)[:self.state_prep_instruction.num_qubits-1]: + qc.x(q) + qc.append(self.state_prep_instruction, list(qc.qubits)[:self.state_prep_instruction.num_qubits]) + return qc + + else: + sub = QuantumCircuit(self.state_prep_instruction.num_qubits, name=label if label is not None else "R_psi") + sub_qubits = list(sub.qubits) + + sub.append(self.state_prep_instruction, sub_qubits) + for q in sub_qubits: + sub.x(q) + sub.h(sub_qubits[-1]) + sub.mcx(sub_qubits[:-1], sub_qubits[-1]) # no-ancilla MCX + sub.h(sub_qubits[-1]) + for q in sub_qubits: + sub.x(q) + sub.append(self.state_prep_instruction.inverse(), sub_qubits) + instr = sub.to_instruction() + qc.append(instr, list(qc.qubits)[:self.state_prep_instruction.num_qubits]) + return instr + + + def build_circuit(self): + for _ in range(self.optimal_iterations): + self.oracle(as_instruction=False) + self.diffuser(as_instruction=False) + diff --git a/solvers/qiskit/knapsack_quantum_tree_generator/QTG/QTG.py b/solvers/qiskit/knapsack_quantum_tree_generator/QTG/QTG.py new file mode 100644 index 00000000..00e806f6 --- /dev/null +++ b/solvers/qiskit/knapsack_quantum_tree_generator/QTG/QTG.py @@ -0,0 +1,185 @@ +import math +from qiskit import QuantumCircuit +from qiskit.circuit import Instruction, Gate +from BaseCircuit.circuit_core import Circuit +from knapsack.knapsack import KnapsackInstance +from math import pi +from typing import Tuple +from qiskit.circuit.library import IntegerComparator, RYGate +from BaseCircuit.adder import constant_qft_add_gate, constant_qft_sub_gate +from BaseCircuit.comparator import apply_threshold_controlled_U + +class QTG(Circuit): + """ + Implements the Quantum Tree Generator (QTG) for generating superpositions of feasible solutions. + """ + + def __init__(self, input_circuit: Circuit, knapsack: KnapsackInstance, name: str = "QTG", depth_interval: Tuple[int, int] = (0, -1), current_best_solution: str= None, bias: float = 0.0, has_ancillas: bool = False): + """ + Initialize the QTG instance. + """ + self.input_circuit = input_circuit + self.name = name + self.knapsack = knapsack + self.current_circuit = self.copy_circuit_structure() + self.bias = bias + self.has_ancillas = has_ancillas + + self.depth_interval = depth_interval + if depth_interval[1] == -1: + self.depth_interval = (depth_interval[0], len(knapsack.items)) + if current_best_solution is None: + self.current_best_solution = '0' * len(knapsack.items) + else: + self.current_best_solution = current_best_solution + + self.initialize_capacity() + print(f"QTG initialized with depth interval {self.depth_interval} and current best solution {self.current_best_solution}") + + def initialize_capacity(self, name_capacity_register: str = "capacity"): + capacity_bin = format(self.knapsack.capacity, f'0{self.knapsack.capacity.bit_length()}b') + # LSB first + reg = self.current_circuit.get_register(name_capacity_register) + for i, bit in enumerate(reversed(capacity_bin)): + if bit == '1': + self.current_circuit.qc.x(reg[i]) + + def copy_circuit_structure(self) -> QuantumCircuit: + qc = Circuit(name=self.name) + for qreg in self.input_circuit.registers.q: + reg = self.input_circuit.get_register(qreg) + qc.add_qubits(reg.name, reg.size) + for areg in self.input_circuit.registers.a: + reg = self.input_circuit.get_register(areg) + qc.add_ancilla(reg.name, reg.size) + for creg in self.input_circuit.registers.c: + reg = self.input_circuit.get_register(creg) + qc.add_clbits(reg.name, reg.size) + return qc + + def apply_U1_ancilla( + self, + m: int, + decisions_reg_name: str = "items", + capacity_reg_name: str = "capacity", + flag_reg_name: str = "oracle", + bias_angle: float = pi / 2 + ): + """ + U1_m: compute feasibility flag 'remaining capacity >= w_m', apply controlled biased rotation on x_m, uncompute flag. + Comparator + biased Hadamard (RY) + uncompute. Flag orientation may be inverted depending on library version. + """ + decision_vars = self.current_circuit.get_register(decisions_reg_name) + capacities = self.current_circuit.get_register(capacity_reg_name) + flag = self.current_circuit.get_register(flag_reg_name) + + + # Comparator: flips flag if cap >= wm + comp = IntegerComparator(num_state_qubits=len(capacities), value=self.knapsack.weights[m], name=f'cap >= w_{m}') + #determine how many ancillas the comparator needs. + need_ancillas = comp.num_qubits - (len(capacities) + 1) + if need_ancillas < 0: + need_ancillas = 0 + anc_slice = [] + if need_ancillas > 0: + anc_pool = None + # Prefer any ancilla register that is not the flag register and has enough qubits + for name, areg in self.current_circuit.registers.a.items(): + if name != flag_reg_name and len(areg) >= need_ancillas: + anc_pool = areg + break + if anc_pool is None: + raise ValueError(f"Insufficient ancillas for comparator: need {need_ancillas}, available " + f"{ {n: len(r) for n, r in self.current_circuit.registers.a.items()} } [1]") + anc_slice = list(anc_pool[:need_ancillas]) + + qargs_cmp = list(capacities) + [flag[0]]+ anc_slice + self.current_circuit.append_instruction(comp, qargs_cmp) + + # Controlled Hadamard on decision qubit x_m + hprime_ctrl = self.biased_hadamard_gate(self.current_best_solution[m]).control(1) + self.current_circuit.append_instruction(hprime_ctrl, [flag[0], decision_vars[m]]) + + # Uncompute the flag to release ancillas back to |0> (reuse same qargs) [1] + self.current_circuit.append_instruction(comp.inverse(), qargs_cmp) + + def apply_U1( + self, + m: int, + decisions_reg_name: str = "items", + capacity_reg_name: str = "capacity", + ): + """ + U1_m: compute feasibility flag 'remaining capacity >= w_m', apply controlled biased rotation on x_m, uncompute flag. + Comparator + biased Hadamard (RY) + uncompute. Flag orientation may be inverted depending on library version. + """ + decision_vars = self.current_circuit.get_register(decisions_reg_name) + capacities = self.current_circuit.get_register(capacity_reg_name) + + U = self.biased_hadamard_gate(self.current_best_solution[m]) + controlled_H = apply_threshold_controlled_U(self.current_circuit.qc, capacities, decision_vars[m], U, self.knapsack.weights[m], inplace=False, label=f"CH(c >= w_{m})") + # Controlled Hadamard on decision qubit x_m + # INPLACE VERSION: + #apply_threshold_controlled_U(self.current_circuit.qc, capacities, decision_vars[m], U, self.knapsack.weights[m], inplace=True) + + self.current_circuit.append_instruction(controlled_H, list(capacities)+ [decision_vars[m]]) + + def apply_U2( + self, + m: int, + decisions_reg_name: str = "items", + capacity_reg_name: str = "capacity", + ): + """ + U2_m: subtract w_m from the capacity register controlled by x_m (QFT subtractor) [1]. + """ + decision_vars = self.current_circuit.get_register(decisions_reg_name) + capacity = self.current_circuit.get_register(capacity_reg_name) + wm = self.knapsack.weights[m] + + sub_gate_ctrl = constant_qft_sub_gate(len(capacity), wm, name=f"Sub(w_{m})").control(1) + self.current_circuit.append_instruction(sub_gate_ctrl, [decision_vars[m]] + list(capacity)) + + def apply_U3( + self, + m: int, + decisions_reg_name: str = "items", + profit_reg_name: str = "profit", + ): + """ + U3_m: add p_m to the profit register controlled by x_m (QFT adder) [1]. + """ + decision_vars = self.current_circuit.get_register(decisions_reg_name) + profit = self.current_circuit.get_register(profit_reg_name) + + add_gate_ctrl = constant_qft_add_gate(len(profit), self.knapsack.values[m], name=f"Add(p_{m})").control(1) + self.current_circuit.append_instruction(add_gate_ctrl, [decision_vars[m]] + list(profit)) + + def build_circuit(self) -> QuantumCircuit: + for i in range(self.depth_interval[0], self.depth_interval[1]): + if self.has_ancillas: + self.apply_U1_ancilla(i) + else: + self.apply_U1(i) + self.apply_U2(i) + self.apply_U3(i) + + def to_instruction(self) -> Instruction: + return self.build_circuit().to_instruction() + + def biased_hadamard_gate(self, current_best_bit: int) -> Gate: + """ + Biased Hadamard via RY(theta) on the decision qubit. + """ + if current_best_bit == '0': + return RYGate(2*math.acos(math.sqrt((1+self.bias)/(2+self.bias)))) + elif current_best_bit == '1': + return RYGate(2*math.acos(math.sqrt((1)/(2+self.bias)))) + + else: + raise ValueError("Invalid current best bit value. Must be '0' or '1'.") + + + + + diff --git a/solvers/qiskit/knapsack_quantum_tree_generator/data/SimulateOptions.py b/solvers/qiskit/knapsack_quantum_tree_generator/data/SimulateOptions.py new file mode 100644 index 00000000..51e372e0 --- /dev/null +++ b/solvers/qiskit/knapsack_quantum_tree_generator/data/SimulateOptions.py @@ -0,0 +1,11 @@ +from dataclasses import dataclass +from typing import Any, Optional + + +@dataclass +class SimulationOptions: + method: str = "statevector" # "statevector" or "automatic" or "density_matrix" (if Aer) + shots: Optional[int] = None + seed_simulator: Optional[int] = None + noise_model: Any = None # from qiskit_aer.noise import NoiseModel + measure_all_if_none: bool = True # For shot-based runs, add measure_all if circuit has none \ No newline at end of file diff --git a/solvers/qiskit/knapsack_quantum_tree_generator/data/TranspileOptions.py b/solvers/qiskit/knapsack_quantum_tree_generator/data/TranspileOptions.py new file mode 100644 index 00000000..fb3fd9a2 --- /dev/null +++ b/solvers/qiskit/knapsack_quantum_tree_generator/data/TranspileOptions.py @@ -0,0 +1,12 @@ +from dataclasses import dataclass +from typing import Any, Optional, Sequence + +@dataclass +class TranspileOptions: + optimization_level: int = 2 + seed_transpiler: Optional[int] = None + basis_gates: Optional[Sequence[str]] = None + layout_method: Optional[str] = None + routing_method: Optional[str] = None + coupling_map: Any = None # can be CouplingMap or list + target: Any = None # qiskit.transpiler.Target (optional) \ No newline at end of file diff --git a/solvers/qiskit/knapsack_quantum_tree_generator/helper.py b/solvers/qiskit/knapsack_quantum_tree_generator/helper.py new file mode 100644 index 00000000..22bffb30 --- /dev/null +++ b/solvers/qiskit/knapsack_quantum_tree_generator/helper.py @@ -0,0 +1,55 @@ +import sys + +from knapsack.knapsack import KnapsackInstance, Item + + +def parse_input(input_data: str): + lines = input_data.strip().split('\n') + if not lines: + raise ValueError("Input data is empty") + + number_items = int(lines[0]) + items = [] + # read items into value, weight lists + # Note: the input format in the issue description has: index value weight + # and the code snippet used indexes[i] in the output. + # We should preserve the original index if we want to return it. + # However, Item class doesn't store original index explicitly other than 'id' which we overwrite. + # Let's add 'original_index' to Item or just use id. + # Actually, the user's snippet does: indexes.append(int(item[0])) + # and then included_indexes.append(indexes[i]) + + # Let's store the original index in the Item object by adding a field or using metadata. + # For now, let's just use the 'id' field to store the original index from the input. + for i in range(number_items): + parts = lines[i + 1].split(' ') + orig_idx = int(parts[0]) + val = int(parts[1]) + weight = int(parts[2]) + items.append(Item(id=orig_idx, value=val, weight=weight)) + + capacity = int(lines[-1]) + + return KnapsackInstance(items=items, capacity=capacity, sort_by_value=False) + + +def run(solve_func): + arg_count = len(sys.argv) - 1 + if arg_count != 2: + raise TypeError( + f'This script expects exactly 2 arguments but got {arg_count}. Input file (argument 1) and output file (argument 2).') + + input_path = sys.argv[1] + output_path = sys.argv[2] + + _run_with_files(input_path, output_path, solve_func) + + +def _run_with_files(input_path: str, output_path: str, solve_func): + with open(input_path, 'r') as f: + input_data = f.read() + + result = solve_func(input_data) + + with open(output_path, 'w') as f: + f.write(result) diff --git a/solvers/qiskit/knapsack_quantum_tree_generator/knapsack/knapsack.py b/solvers/qiskit/knapsack_quantum_tree_generator/knapsack/knapsack.py new file mode 100644 index 00000000..e563aaeb --- /dev/null +++ b/solvers/qiskit/knapsack_quantum_tree_generator/knapsack/knapsack.py @@ -0,0 +1,95 @@ +import yaml # type: ignore +from dataclasses import dataclass +from pathlib import Path +from typing import Any, Dict + +@dataclass +class Item: + id: int + weight: int + value: int + +class KnapsackInstance: + """ + Items are sorted by value after loading by default. + """ + def __init__(self, config_path: str | Path | None = None, items: list[Item] | None = None, capacity: int | None = None, sort_by_value: bool = True): + if config_path: + with open(config_path, 'r', encoding='utf-8') as f: + cfg = yaml.safe_load(f) + if cfg is None: + raise ValueError(f"Configuration file at {config_path} invalid.") + + self.metadata = cfg.get('metadata', {}) + self.capacity = int(cfg.get('capacity', None)) + self.items = [ + Item( + id=None, + weight=int(row['weight']), + value=int(row['value']) + ) + for i, row in enumerate(cfg.get('items', [])) + ] + elif items is not None and capacity is not None: + self.metadata = {} + self.capacity = capacity + self.items = items + else: + raise ValueError("Either config_path or (items and capacity) must be provided.") + + if sort_by_value: + self.sort_items_by_value() + self.assign_ids_by_rank() + self.num_items = len(self.items) + self.ids = [it.id for it in self.items] + self.weights = [it.weight for it in self.items] + self.values = [it.value for it in self.items] + + + def get_item_by_id(self, item_id: int): + """Return item with given id, or None if not found.""" + return next((it for it in self.items if it.id == item_id), None) + + def get_properties_dict(self) -> Dict[str, Any]: + """Return instance properties as a dictionary.""" + return { + 'metadata': self.metadata, + 'capacity': self.capacity, + 'items': [it.__dict__ for it in self.items] + } + + def assign_ids_by_rank(self) -> None: + """Assign item ids based on current order in self.items if id is None. Used to sort by density""" + for new_id, it in enumerate(self.items): + if it.id is None: + it.id = new_id + + def sort_items_by_density(self) -> None: + """Sort items in-place by value/weight density in descending order.""" + self.items.sort(key=lambda it: it.value / it.weight if it.weight > 0 else 0, reverse=True) + + def sort_items_by_value(self) -> None: + """Sort items in-place by value in descending order.""" + self.items.sort(key=lambda it: it.value, reverse=True) + + def is_sorted_by_density(self) -> bool: + """Check if items are sorted by value/weight density in descending order.""" + return all((self.items[i].value / self.items[i].weight if self.items[i].weight > 0 else 0) >= + (self.items[i + 1].value / self.items[i + 1].weight if self.items[i + 1].weight > 0 else 0) + for i in range(len(self.items) - 1)) + + def is_sorted_by_value(self) -> bool: + """Check if items are sorted by value in descending order.""" + return all(self.items[i].value >= self.items[i + 1].value for i in range(len(self.items) - 1)) + + def _copy(self) -> 'KnapsackInstance': + """Create a deep copy of the KnapsackInstance.""" + new_instance = KnapsackInstance.__new__(KnapsackInstance) + new_instance.metadata = self.metadata.copy() + new_instance.capacity = self.capacity + new_instance.items = [Item(id=it.id, weight=it.weight, value=it.value) for it in self.items] + new_instance.num_items = self.num_items + new_instance.ids = self.ids.copy() + new_instance.weights = self.weights.copy() + new_instance.values = self.values.copy() + return new_instance diff --git a/solvers/qiskit/knapsack_quantum_tree_generator/knapsack_quantum_tree_generator_openqasm.py b/solvers/qiskit/knapsack_quantum_tree_generator/knapsack_quantum_tree_generator_openqasm.py new file mode 100644 index 00000000..2d2a98f6 --- /dev/null +++ b/solvers/qiskit/knapsack_quantum_tree_generator/knapsack_quantum_tree_generator_openqasm.py @@ -0,0 +1,45 @@ +from qiskit import qasm2 + +from BaseCircuit.circuit_core import Circuit +from QTG.QTG import QTG +from helper import parse_input +from knapsack.knapsack import KnapsackInstance +from helper import run + + +def _retrieve_openqasm(knapsack: KnapsackInstance): + # Prepare the quantum circuit + circuit = Circuit("Knapsack_Demo") + has_ancillas = False + circuit.prepare_knapsack_circuit(knapsack, has_ancillas=has_ancillas) + + # Build the QTG + bias = 0.5 + current_best_solution = "0" * len(knapsack.items) + + qtg = QTG( + input_circuit=circuit, + knapsack=knapsack, + depth_interval=(0, -1), + bias=bias, + current_best_solution=current_best_solution, + has_ancillas=has_ancillas + ) + qtg.build_circuit() + + # Append QTG + qreg = circuit.get_qubits_in_registers(all=True) + circuit.append_subcircuit_as_instruction(qtg.current_circuit, qubits=qreg, + name='qtg_circuit') + + circuit.measure_items("items_c") + + return qasm2.dumps(circuit.qc) + + +# def _retrieve_openqasm(input_data): +# return _run_retrieve_openqasm(parse_input(input_data)) + + +if __name__ == "__main__": + run(lambda input_data: _retrieve_openqasm(parse_input(input_data))) diff --git a/solvers/qiskit/knapsack_quantum_tree_generator/knapsack_quantum_tree_generator_solve.py b/solvers/qiskit/knapsack_quantum_tree_generator/knapsack_quantum_tree_generator_solve.py new file mode 100644 index 00000000..e1e8632f --- /dev/null +++ b/solvers/qiskit/knapsack_quantum_tree_generator/knapsack_quantum_tree_generator_solve.py @@ -0,0 +1,93 @@ +from BaseCircuit.circuit_core import Circuit +from BaseCircuit.simulator import Simulator +from QTG.QTG import QTG +from data.SimulateOptions import SimulationOptions +from helper import parse_input +from knapsack.knapsack import KnapsackInstance +from helper import run + + +def _create_output(solve_results): + probabilities = solve_results["probabilities"] + knapsack = solve_results["knapsack"] + + # Best feasible + best_feasible_bitstring = None + best_value = -1 + + for bitstring, _ in probabilities.items(): + weight = sum( + knapsack.items[i].weight for i, b in enumerate(bitstring) if + b == '1') + if weight <= knapsack.capacity: + value = sum( + knapsack.items[i].value for i, b in enumerate(bitstring) if + b == '1') + if best_value == -1 or value > best_value: + best_value = value + best_feasible_bitstring = bitstring + + if best_feasible_bitstring: + included_indexes = [] + for i, bit in enumerate(best_feasible_bitstring): + if bit == '1': + included_indexes.append(knapsack.items[i].id) + + # Sort indexes to match expected output if necessary, + # but the original code just appended them. + return f"{best_value}\n{included_indexes}" + else: + return "0\n[]" + + +def _get_solution(knapsack: KnapsackInstance): + # Prepare the quantum circuit + circuit = Circuit("Knapsack_Demo") + has_ancillas = False + circuit.prepare_knapsack_circuit(knapsack, has_ancillas=has_ancillas) + + # Build the QTG + bias = 0.5 + current_best_solution = "0" * len(knapsack.items) + + qtg = QTG( + input_circuit=circuit, + knapsack=knapsack, + depth_interval=(0, -1), + bias=bias, + current_best_solution=current_best_solution, + has_ancillas=has_ancillas + ) + qtg.build_circuit() + + # Append QTG + qreg = circuit.get_qubits_in_registers(all=True) + circuit.append_subcircuit_as_instruction(qtg.current_circuit, qubits=qreg, + name='qtg_circuit') + + print(circuit.draw()) + # Add measurements + circuit.measure_items("items_c") + + # Simulate + sim_opts = SimulationOptions(method="automatic", shots=50000) + simulator = Simulator(circuit) + simulator.simulate(sim_opts) + + # Process results + probabilities = simulator.get_probabilities_from_counts() + + return { + "probabilities": probabilities, + "knapsack": knapsack + } + + +def _solve(input_data): + solution = _get_solution(parse_input(input_data)) + print(solution) + return _create_output(solution) + + +if __name__ == "__main__": + run(_solve) diff --git a/solvers/qiskit/knapsack_quantum_tree_generator/requirements.txt b/solvers/qiskit/knapsack_quantum_tree_generator/requirements.txt new file mode 100644 index 00000000..d48a4d37 --- /dev/null +++ b/solvers/qiskit/knapsack_quantum_tree_generator/requirements.txt @@ -0,0 +1,24 @@ +contourpy==1.3.3 +cycler==0.12.1 +dill==0.4.1 +docplex==2.32.264 +fonttools==4.62.1 +kiwisolver==1.5.0 +matplotlib==3.10.9 +networkx==3.6.1 +numpy==2.4.4 +packaging==26.2 +pillow==12.2.0 +psutil==7.2.2 +pyparsing==3.3.2 +python-dateutil==2.9.0.post0 +PyYAML==6.0.3 +qiskit==2.4.1 +qiskit-aer==0.17.2 +qiskit-optimization==0.7.0 +rustworkx==0.17.1 +scipy==1.17.1 +setuptools==82.0.1 +six==1.17.0 +stevedore==5.7.0 +typing_extensions==4.15.0 diff --git a/solvers/tools/equivalencechecking/equivalencechecking.py b/solvers/tools/equivalencechecking/equivalencechecking.py new file mode 100644 index 00000000..4cb20737 --- /dev/null +++ b/solvers/tools/equivalencechecking/equivalencechecking.py @@ -0,0 +1,39 @@ +import json +import sys +from pathlib import Path +from typing import Any + +from run import run + + +def _read_json(path: Path) -> dict[str, Any]: + return json.loads(path.read_text(encoding="utf-8")) + + +def _write_json(path: Path, data: dict[str, Any]) -> None: + path.write_text(json.dumps(data, indent=2), encoding="utf-8") + + +def main() -> int: + if len(sys.argv) != 3: + print("Usage: api_file.py ", + file=sys.stderr) + return 2 + + input_path = Path(sys.argv[1]) + output_path = Path(sys.argv[2]) + payload = _read_json(input_path) + + result = run( + strategy=payload["strategy"], + qasm_a=payload["qasmA"], + qasm_b=payload["qasmB"], + qcec_options=payload.get("qcecOptions", {}), + ) + _write_json(output_path, result) + + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/solvers/tools/equivalencechecking/requirements.txt b/solvers/tools/equivalencechecking/requirements.txt new file mode 100644 index 00000000..89a5e22a --- /dev/null +++ b/solvers/tools/equivalencechecking/requirements.txt @@ -0,0 +1,47 @@ +annotated-doc==0.0.4 +annotated-types==0.7.0 +anyio==4.14.1 +asttokens==3.0.1 +click==8.4.2 +comm==0.2.3 +decorator==5.3.1 +executing==2.2.1 +fastapi==0.139.0 +h11==0.16.0 +idna==3.18 +iniconfig==2.3.0 +ipython==9.15.0 +ipython_pygments_lexers==1.1.1 +ipywidgets==8.1.8 +jedi==0.20.0 +jupyterlab_widgets==3.0.16 +lark==1.3.1 +matplotlib-inline==0.2.2 +mqt-core==3.6.1 +mqt.qcec==3.6.1 +numpy==2.4.6 +packaging==26.0 +parso==0.8.7 +pexpect==4.9.0 +pluggy==1.6.0 +prompt_toolkit==3.0.52 +psutil==7.2.2 +ptyprocess==0.7.0 +pure_eval==0.2.3 +pydantic==2.13.4 +pydantic_core==2.46.4 +Pygments==2.20.0 +pyperclip==1.11.0 +pytest==9.0.3 +pyzx==0.10.4 +setuptools==82.0.1 +stack-data==0.6.3 +starlette==1.3.1 +tqdm==4.68.4 +traitlets==5.15.1 +typing-inspection==0.4.2 +typing_extensions==4.16.0 +uvicorn==0.51.0 +wcwidth==0.8.2 +wheel==0.46.3 +widgetsnbextension==4.0.15 diff --git a/solvers/tools/equivalencechecking/run.py b/solvers/tools/equivalencechecking/run.py new file mode 100644 index 00000000..5b4ae09e --- /dev/null +++ b/solvers/tools/equivalencechecking/run.py @@ -0,0 +1,21 @@ +from typing import Any + +from run_pyzx import run_pyzx +from run_qcec import run_qcec_isolated + + +def run(strategy: str, qasm_a: str, qasm_b: str, + qcec_options: dict[str, Any] | None = None) -> dict[str, Any]: + if strategy == "pyzx": + return run_pyzx( + qasm_a=qasm_a, + qasm_b=qasm_b, + ) + elif strategy == "mqt-qcec": + return run_qcec_isolated( + qasm_a=qasm_a, + qasm_b=qasm_b, + qcec_options=qcec_options, + ) + else: + raise ValueError(f"Unknown strategy: {strategy}") diff --git a/solvers/tools/equivalencechecking/run_pyzx.py b/solvers/tools/equivalencechecking/run_pyzx.py new file mode 100644 index 00000000..60cbdc8f --- /dev/null +++ b/solvers/tools/equivalencechecking/run_pyzx.py @@ -0,0 +1,52 @@ +#!/usr/bin/env python3 +import time +from pathlib import Path +from typing import Any + +import pyzx as zx + + +def run_pyzx( + qasm_a: str, + qasm_b: str, +) -> dict[str, Any]: + started = time.perf_counter() + + try: + circuit_a = zx.Circuit.from_qasm(qasm_a) + circuit_b = zx.Circuit.from_qasm(qasm_b) + + equivalent = circuit_a.verify_equality(circuit_b) + + runtime_ms = round((time.perf_counter() - started) * 1000, 3) + + return { + "strategy": "pyzx", + # False vorsichtshalber nicht als bewiesene Nichtäquivalenz behandeln. + "status": "equivalent" if equivalent else "unknown", + "globalPhaseIgnored": True if equivalent else None, + "runtimeMs": runtime_ms, + "rawEquivalence": str(equivalent), + "message": ( + None + if equivalent + else "PyZX konnte die Äquivalenz nicht beweisen." + ), + "error": None, + } + + except Exception as exc: + runtime_ms = round((time.perf_counter() - started) * 1000, 3) + + return { + "strategy": "pyzx", + "status": "error", + "globalPhaseIgnored": None, + "runtimeMs": runtime_ms, + "rawEquivalence": None, + "message": str(exc), + "error": { + "type": type(exc).__name__, + "message": str(exc), + }, + } diff --git a/solvers/tools/equivalencechecking/run_qcec.py b/solvers/tools/equivalencechecking/run_qcec.py new file mode 100644 index 00000000..e1e7b86d --- /dev/null +++ b/solvers/tools/equivalencechecking/run_qcec.py @@ -0,0 +1,176 @@ +#!/usr/bin/env python3 +import multiprocessing +import os +import time +from typing import Any + +from mqt import qcec + +DEFAULT_QCEC_OPTIONS: dict[str, Any] = { + # QCEC's parallel checker can remain stuck while joining its native worker + # threads. + "parallel": False, +} +DEFAULT_HARD_TIMEOUT_SECONDS = float( + os.getenv("QCEC_HARD_TIMEOUT_SECONDS", "10") +) + + +def normalize_status(raw_equivalence: str) -> tuple[str, bool | None]: + value = raw_equivalence.lower() + + if "not" in value and "equivalent" in value: + return "not_equivalent", None + + if "up_to_global_phase" in value or "global_phase" in value: + return "equivalent", True + + if "equivalent" in value: + return "equivalent", False + + if "unknown" in value or "no_information" in value or "inconclusive" in value: + return "unknown", None + + return "unknown", None + + +def stringify_equivalence(result: Any) -> str: + equivalence = getattr(result, "equivalence", result) + + if callable(equivalence): + equivalence = equivalence() + + if hasattr(equivalence, "name"): + return str(equivalence.name) + + if hasattr(equivalence, "value"): + return str(equivalence.value) + + return str(equivalence) + + +def run_qcec( + qasm_a: str, + qasm_b: str, + qcec_options: dict[str, Any] | None = None, +) -> dict[str, Any]: + started = time.perf_counter() + effective_options = DEFAULT_QCEC_OPTIONS | (qcec_options or {}) + + try: + result = qcec.verify(qasm_a, qasm_b, **effective_options) + + raw_equivalence = stringify_equivalence(result) + status, global_phase_ignored = normalize_status(raw_equivalence) + + runtime_ms = round((time.perf_counter() - started) * 1000, 3) + + return { + "strategy": "mqt-qcec", + "status": status, + "globalPhaseIgnored": global_phase_ignored, + "runtimeMs": runtime_ms, + "rawEquivalence": raw_equivalence, + "message": None, + "error": None, + } + + except Exception as exc: + runtime_ms = round((time.perf_counter() - started) * 1000, 3) + + return { + "strategy": "mqt-qcec", + "status": "error", + "globalPhaseIgnored": None, + "runtimeMs": runtime_ms, + "rawEquivalence": None, + "message": str(exc), + "error": { + "type": type(exc).__name__, + "message": str(exc), + }, + } + + +def _isolated_qcec_worker( + connection: Any, + qasm_a: str, + qasm_b: str, + qcec_options: dict[str, Any] | None, +) -> None: + try: + connection.send(run_qcec(qasm_a, qasm_b, qcec_options)) + finally: + connection.close() + + +def run_qcec_isolated( + qasm_a: str, + qasm_b: str, + qcec_options: dict[str, Any] | None = None, + hard_timeout_seconds: float = DEFAULT_HARD_TIMEOUT_SECONDS, +) -> dict[str, Any]: + """Run QCEC in a disposable process with a reliable wall-clock timeout.""" + if hard_timeout_seconds <= 0: + raise ValueError("hard_timeout_seconds must be greater than zero") + + started = time.perf_counter() + context = multiprocessing.get_context("spawn") + receiving_connection, sending_connection = context.Pipe(duplex=False) + process = context.Process( + target=_isolated_qcec_worker, + args=(sending_connection, qasm_a, qasm_b, qcec_options), + daemon=True, + ) + + try: + process.start() + sending_connection.close() + + if receiving_connection.poll(hard_timeout_seconds): + try: + result = receiving_connection.recv() + except EOFError: + result = None + + process.join(timeout=1) + if result is not None: + return result + + message = ( + "The QCEC worker exited without returning a result " + f"(exit code {process.exitcode})." + ) + error_type = "WorkerProcessError" + else: + message = ( + "MQT QCEC exceeded the hard timeout of " + f"{hard_timeout_seconds:g} seconds." + ) + error_type = "TimeoutError" + except Exception as exc: + message = str(exc) + error_type = type(exc).__name__ + finally: + receiving_connection.close() + sending_connection.close() + if process.is_alive(): + process.terminate() + process.join(timeout=1) + if process.is_alive(): + process.kill() + process.join() + + runtime_ms = round((time.perf_counter() - started) * 1000, 3) + return { + "strategy": "mqt-qcec", + "status": "error", + "globalPhaseIgnored": None, + "runtimeMs": runtime_ms, + "rawEquivalence": None, + "message": message, + "error": { + "type": error_type, + "message": message, + }, + } diff --git a/src/main/java/edu/kit/provideq/toolbox/api/tools/EquivalenceCheckingRouter.java b/src/main/java/edu/kit/provideq/toolbox/api/tools/EquivalenceCheckingRouter.java new file mode 100644 index 00000000..53b07a65 --- /dev/null +++ b/src/main/java/edu/kit/provideq/toolbox/api/tools/EquivalenceCheckingRouter.java @@ -0,0 +1,65 @@ +package edu.kit.provideq.toolbox.api.tools; + +import static org.springdoc.webflux.core.fn.SpringdocRouteBuilder.route; +import static org.springframework.http.MediaType.APPLICATION_JSON; +import static org.springframework.web.reactive.function.server.RequestPredicates.accept; +import static org.springframework.web.reactive.function.server.RequestPredicates.contentType; +import static org.springframework.web.reactive.function.server.ServerResponse.ok; + +import com.fasterxml.jackson.databind.JsonNode; +import edu.kit.provideq.toolbox.tools.equivalencechecking.EquivalenceChecking; +import org.springframework.context.annotation.Bean; +import org.springframework.context.annotation.Configuration; +import org.springframework.http.HttpStatus; +import org.springframework.web.reactive.config.EnableWebFlux; +import org.springframework.web.reactive.function.server.RouterFunction; +import org.springframework.web.reactive.function.server.ServerRequest; +import org.springframework.web.reactive.function.server.ServerResponse; +import org.springframework.web.server.ResponseStatusException; +import reactor.core.publisher.Mono; +import reactor.core.scheduler.Schedulers; + +/** Routes requests to synchronous toolbox utilities. */ +@Configuration +@EnableWebFlux +public class EquivalenceCheckingRouter { + static final String EQUIVALENCE_CHECKING_PATH = "/tools/equivalencechecking"; + + private final EquivalenceChecking equivalenceChecking; + + public EquivalenceCheckingRouter(EquivalenceChecking equivalenceChecking) { + this.equivalenceChecking = equivalenceChecking; + } + + /** Registers the equivalence-checking endpoint. */ + @Bean + public RouterFunction getEquivalenceCheckingRoutes() { + return route().POST( + EQUIVALENCE_CHECKING_PATH, + contentType(APPLICATION_JSON).and(accept(APPLICATION_JSON)), + this::handleEquivalenceChecking, + ops -> ops + .operationId("equivalenceChecking") + .tag("tools") + .description("Checks whether two quantum circuits are equivalent.") + ).build(); + } + + private Mono handleEquivalenceChecking(ServerRequest request) { + return request.bodyToMono(JsonNode.class) + .switchIfEmpty(Mono.error(new ResponseStatusException( + HttpStatus.BAD_REQUEST, + "A JSON request body is required."))) + .flatMap(input -> Mono.fromCallable(() -> equivalenceChecking.check(input)) + .subscribeOn(Schedulers.boundedElastic())) + .onErrorMap( + IllegalStateException.class, + exception -> new ResponseStatusException( + HttpStatus.INTERNAL_SERVER_ERROR, + exception.getMessage(), + exception)) + .flatMap(output -> ok() + .contentType(APPLICATION_JSON) + .bodyValue(output)); + } +} diff --git a/src/main/java/edu/kit/provideq/toolbox/circuit/processing/CircuitProcessingConfiguration.java b/src/main/java/edu/kit/provideq/toolbox/circuit/processing/CircuitProcessingConfiguration.java new file mode 100644 index 00000000..53379736 --- /dev/null +++ b/src/main/java/edu/kit/provideq/toolbox/circuit/processing/CircuitProcessingConfiguration.java @@ -0,0 +1,63 @@ +package edu.kit.provideq.toolbox.circuit.processing; + +import edu.kit.provideq.toolbox.ResourceProvider; +import edu.kit.provideq.toolbox.circuit.processing.results.Result; +import edu.kit.provideq.toolbox.circuit.processing.solver.MoveToExecutionSolver; +import edu.kit.provideq.toolbox.circuit.processing.solver.MoveToMitigationSolver; +import edu.kit.provideq.toolbox.circuit.processing.solver.MoveToOptimizationSolver; +import edu.kit.provideq.toolbox.exception.MissingExampleException; +import edu.kit.provideq.toolbox.meta.Problem; +import edu.kit.provideq.toolbox.meta.ProblemManager; +import edu.kit.provideq.toolbox.meta.ProblemType; +import java.io.IOException; +import java.util.HashSet; +import java.util.Objects; +import java.util.Set; +import org.springframework.context.annotation.Bean; +import org.springframework.context.annotation.Configuration; + +@Configuration +public class CircuitProcessingConfiguration { + public static final ProblemType CIRCUIT_PROCESSING = new ProblemType<>( + "circuit-processing", + "A quantum circuit processing problem that routes a QASM circuit through optimization, " + + "error mitigation, or execution.", + String.class, + Result.class + ); + + @Bean + ProblemManager getCircuitProcessingManager( + ResourceProvider provider, + MoveToExecutionSolver moveToExecutionSolver, + MoveToOptimizationSolver moveToOptimizationSolver, + MoveToMitigationSolver moveToMitigationSolver + ) { + return new ProblemManager<>( + CIRCUIT_PROCESSING, + Set.of( + moveToExecutionSolver, + moveToOptimizationSolver, + moveToMitigationSolver + ), + loadExampleProblems(provider) + ); + } + + private Set> loadExampleProblems(ResourceProvider provider) { + try { + String[] problemNames = new String[] {"bell-state.qasm", "cswap.qasm"}; + var problemSet = new HashSet>(); + for (var problemName : problemNames) { + var problemStream = Objects.requireNonNull( + getClass().getResourceAsStream(problemName), "Problem " + problemName + " not found"); + var problem = new Problem<>(CIRCUIT_PROCESSING); + problem.setInput(provider.readStream(problemStream)); + problemSet.add(problem); + } + return problemSet; + } catch (IOException e) { + throw new MissingExampleException(CIRCUIT_PROCESSING, e); + } + } +} diff --git a/src/main/java/edu/kit/provideq/toolbox/circuit/processing/results/ExecutionResult.java b/src/main/java/edu/kit/provideq/toolbox/circuit/processing/results/ExecutionResult.java new file mode 100644 index 00000000..30491be0 --- /dev/null +++ b/src/main/java/edu/kit/provideq/toolbox/circuit/processing/results/ExecutionResult.java @@ -0,0 +1,79 @@ +package edu.kit.provideq.toolbox.circuit.processing.results; + +import edu.kit.provideq.toolbox.util.Pair; +import java.util.List; +import java.util.Optional; + +/** + * A record representing the result of a circuit processing execution operation. + * This includes information about measurement outcomes, probabilities, + * the resultant circuit, and a generated result string. + * + * @param sortedMeasurements An optional list of measurement results sorted by frequency in descending order. + * Each entry consists of a string representing the measurement outcome + * and an integer representing the count of that outcome. + * @param sortedProbabilities An optional list of measurement probabilities sorted in descending order. + * Each entry consists of a string representing the measurement outcome + * and a double representing the probability of that outcome. + * @param resultString An optional string representation of the circuit processing result. + * @param circuit An optional string representing the processed circuit in a predefined format. + */ +public record ExecutionResult( + Optional>> sortedMeasurements, + Optional>> sortedProbabilities, + Optional resultString, + Optional circuit +) implements Result { + + @Override + public R accept(ResultVisitor resultVisitor) { + return resultVisitor.visit(this); + } + + public boolean hasResult() { + return resultString.isPresent(); + } + + public boolean hasCircuit() { + return circuit.isPresent(); + } + + public Optional getHighestResult() { + if (sortedMeasurements().isEmpty()) { + return Optional.empty(); + } + + var sortedMeasurements = sortedMeasurements().get(); + if (sortedMeasurements.isEmpty()) { + return Optional.empty(); + } + + return Optional.of(sortedMeasurements.get(0).first()); + } + + public Optional getHighestResultMeasurement() { + if (sortedMeasurements().isEmpty()) { + return Optional.empty(); + } + + var sortedMeasurements = sortedMeasurements().get(); + if (sortedMeasurements.isEmpty()) { + return Optional.empty(); + } + + return Optional.of(sortedMeasurements.get(0).second()); + } + + public Optional getHighestResultProbability() { + if (sortedProbabilities().isEmpty()) { + return Optional.empty(); + } + + var sortedProbabilities = sortedProbabilities().get(); + if (sortedProbabilities.isEmpty()) { + return Optional.empty(); + } + + return Optional.of(sortedProbabilities.get(0).second()); + } +} diff --git a/src/main/java/edu/kit/provideq/toolbox/circuit/processing/results/ExecutionResultHelper.java b/src/main/java/edu/kit/provideq/toolbox/circuit/processing/results/ExecutionResultHelper.java new file mode 100644 index 00000000..c39769dd --- /dev/null +++ b/src/main/java/edu/kit/provideq/toolbox/circuit/processing/results/ExecutionResultHelper.java @@ -0,0 +1,87 @@ +package edu.kit.provideq.toolbox.circuit.processing.results; + +import edu.kit.provideq.toolbox.util.Pair; +import java.util.HashMap; +import java.util.List; +import java.util.Map; +import java.util.Optional; +import java.util.regex.Matcher; +import java.util.regex.Pattern; + +public final class ExecutionResultHelper { + private ExecutionResultHelper() { + throw new IllegalStateException("Utility class"); + } + + @SuppressWarnings("OptionalUsedAsFieldOrParameterType") + public static ExecutionResult createExecutionResult( + Optional resultStringOptional, + Optional circuitOptional + ) { + if (resultStringOptional.isEmpty()) { + return new ExecutionResult( + Optional.empty(), + Optional.empty(), + resultStringOptional, + circuitOptional + ); + } + + var resultString = resultStringOptional.get(); + var counts = parseCounter(resultString); + var sortedMeasurements = createSortedMeasurements(counts); + var sortedProbabilities = createSortedProbabilities(counts); + + return new ExecutionResult( + Optional.of(sortedMeasurements), + Optional.of(sortedProbabilities), + resultStringOptional, + circuitOptional + ); + } + + private static Map parseCounter(String counterString) { + Pattern pattern = Pattern.compile("\\(([01](?:\\s*,\\s*[01])*)\\)\\s*:\\s*(\\d+)"); + + Matcher matcher = pattern.matcher(counterString); + + Map counts = new HashMap<>(); + + while (matcher.find()) { + String bitString = matcher.group(1).replaceAll("\\s*,\\s*", ""); + int count = Integer.parseInt(matcher.group(2)); + + counts.put(bitString, count); + } + + return counts; + } + + private static List> createSortedMeasurements( + Map counts + ) { + return counts.entrySet() + .stream() + .sorted(Map.Entry.comparingByValue().reversed()) + .map(entry -> new Pair<>(entry.getKey(), entry.getValue())) + .toList(); + } + + private static List> createSortedProbabilities( + Map counts + ) { + int totalShots = counts.values() + .stream() + .mapToInt(Integer::intValue) + .sum(); + + return counts.entrySet() + .stream() + .sorted(Map.Entry.comparingByValue().reversed()) + .map(entry -> new Pair<>( + entry.getKey(), + (double) entry.getValue() / totalShots + )) + .toList(); + } +} diff --git a/src/main/java/edu/kit/provideq/toolbox/circuit/processing/results/ExecutionResultVisitor.java b/src/main/java/edu/kit/provideq/toolbox/circuit/processing/results/ExecutionResultVisitor.java new file mode 100644 index 00000000..f0701267 --- /dev/null +++ b/src/main/java/edu/kit/provideq/toolbox/circuit/processing/results/ExecutionResultVisitor.java @@ -0,0 +1,13 @@ +package edu.kit.provideq.toolbox.circuit.processing.results; + +public class ExecutionResultVisitor implements ResultVisitor { + @Override + public String visit(StringResult stringResult) { + return stringResult.value(); + } + + @Override + public String visit(ExecutionResult executionResult) { + return executionResult.getHighestResult().orElse(""); + } +} diff --git a/src/main/java/edu/kit/provideq/toolbox/circuit/processing/results/Result.java b/src/main/java/edu/kit/provideq/toolbox/circuit/processing/results/Result.java new file mode 100644 index 00000000..f1812fff --- /dev/null +++ b/src/main/java/edu/kit/provideq/toolbox/circuit/processing/results/Result.java @@ -0,0 +1,5 @@ +package edu.kit.provideq.toolbox.circuit.processing.results; + +public interface Result { + R accept(ResultVisitor resultVisitor); +} diff --git a/src/main/java/edu/kit/provideq/toolbox/circuit/processing/results/ResultVisitor.java b/src/main/java/edu/kit/provideq/toolbox/circuit/processing/results/ResultVisitor.java new file mode 100644 index 00000000..19a8d7db --- /dev/null +++ b/src/main/java/edu/kit/provideq/toolbox/circuit/processing/results/ResultVisitor.java @@ -0,0 +1,7 @@ +package edu.kit.provideq.toolbox.circuit.processing.results; + +public interface ResultVisitor { + R visit(StringResult stringResult); + + R visit(ExecutionResult executionResult); +} diff --git a/src/main/java/edu/kit/provideq/toolbox/circuit/processing/results/StringResult.java b/src/main/java/edu/kit/provideq/toolbox/circuit/processing/results/StringResult.java new file mode 100644 index 00000000..6d8d8304 --- /dev/null +++ b/src/main/java/edu/kit/provideq/toolbox/circuit/processing/results/StringResult.java @@ -0,0 +1,10 @@ +package edu.kit.provideq.toolbox.circuit.processing.results; + +public record StringResult( + String value +) implements Result { + @Override + public R accept(ResultVisitor resultVisitor) { + return resultVisitor.visit(this); + } +} diff --git a/src/main/java/edu/kit/provideq/toolbox/circuit/processing/solver/CircuitProcessingSolver.java b/src/main/java/edu/kit/provideq/toolbox/circuit/processing/solver/CircuitProcessingSolver.java new file mode 100644 index 00000000..6b192b23 --- /dev/null +++ b/src/main/java/edu/kit/provideq/toolbox/circuit/processing/solver/CircuitProcessingSolver.java @@ -0,0 +1,21 @@ +package edu.kit.provideq.toolbox.circuit.processing.solver; + +import edu.kit.provideq.toolbox.circuit.processing.CircuitProcessingConfiguration; +import edu.kit.provideq.toolbox.circuit.processing.results.Result; +import edu.kit.provideq.toolbox.meta.ProblemSolver; +import edu.kit.provideq.toolbox.meta.ProblemType; +import edu.kit.provideq.toolbox.meta.SubRoutineDefinition; + +public abstract class CircuitProcessingSolver implements ProblemSolver { + public static final SubRoutineDefinition CIRCUIT_PROCESSING_SUBROUTINE = + new SubRoutineDefinition<>( + CircuitProcessingConfiguration.CIRCUIT_PROCESSING, + "Creates a circuit processing solver", + true + ); + + @Override + public ProblemType getProblemType() { + return CircuitProcessingConfiguration.CIRCUIT_PROCESSING; + } +} diff --git a/src/main/java/edu/kit/provideq/toolbox/circuit/processing/solver/MoveToExecutionSolver.java b/src/main/java/edu/kit/provideq/toolbox/circuit/processing/solver/MoveToExecutionSolver.java new file mode 100644 index 00000000..71954d61 --- /dev/null +++ b/src/main/java/edu/kit/provideq/toolbox/circuit/processing/solver/MoveToExecutionSolver.java @@ -0,0 +1,58 @@ +package edu.kit.provideq.toolbox.circuit.processing.solver; + +import edu.kit.provideq.toolbox.Solution; +import edu.kit.provideq.toolbox.SolutionStatus; +import edu.kit.provideq.toolbox.circuit.processing.results.Result; +import edu.kit.provideq.toolbox.circuit.processing.solver.executor.ExecutorConfiguration; +import edu.kit.provideq.toolbox.meta.SolvingProperties; +import edu.kit.provideq.toolbox.meta.SubRoutineDefinition; +import edu.kit.provideq.toolbox.meta.SubRoutineResolver; +import java.util.List; +import org.springframework.stereotype.Component; +import reactor.core.publisher.Mono; + +@Component +public class MoveToExecutionSolver extends CircuitProcessingSolver { + private static final SubRoutineDefinition EXECUTOR_SUBROUTINE = + new SubRoutineDefinition<>( + ExecutorConfiguration.EXECUTOR_CONFIG, + "Creates a execution solver", + true + ); + + @Override + public String getName() { + return "Execute QASM Code"; + } + + @Override + public String getDescription() { + return "Move QASM input to the executors"; + } + + @Override + public List> getSubRoutines() { + return List.of(EXECUTOR_SUBROUTINE); + } + + @Override + public Mono> solve( + String input, + SubRoutineResolver subRoutineResolver, + SolvingProperties properties + ) { + return subRoutineResolver.runSubRoutine(EXECUTOR_SUBROUTINE, input) + .map(executionResultSolution -> { + Solution solution = new Solution<>(this); + SolutionStatus status = executionResultSolution.getStatus(); + if (status == SolutionStatus.ERROR) { + solution.fail(); + solution.setDebugData(executionResultSolution.getDebugData()); + return solution; + } + solution.complete(); + solution.setSolutionData(executionResultSolution.getSolutionData()); + return solution; + }); + } +} diff --git a/src/main/java/edu/kit/provideq/toolbox/circuit/processing/solver/MoveToMitigationSolver.java b/src/main/java/edu/kit/provideq/toolbox/circuit/processing/solver/MoveToMitigationSolver.java new file mode 100644 index 00000000..3da35e5e --- /dev/null +++ b/src/main/java/edu/kit/provideq/toolbox/circuit/processing/solver/MoveToMitigationSolver.java @@ -0,0 +1,45 @@ +package edu.kit.provideq.toolbox.circuit.processing.solver; + +import edu.kit.provideq.toolbox.Solution; +import edu.kit.provideq.toolbox.circuit.processing.results.Result; +import edu.kit.provideq.toolbox.circuit.processing.solver.mitigation.ErrorMitigationConfiguration; +import edu.kit.provideq.toolbox.meta.SolvingProperties; +import edu.kit.provideq.toolbox.meta.SubRoutineDefinition; +import edu.kit.provideq.toolbox.meta.SubRoutineResolver; +import java.util.List; +import org.springframework.stereotype.Component; +import reactor.core.publisher.Mono; + +@Component +public class MoveToMitigationSolver extends CircuitProcessingSolver { + private static final SubRoutineDefinition MITIGATOR_SUBROUTINE = + new SubRoutineDefinition<>( + ErrorMitigationConfiguration.MITIGATION_CONFIG, + "Creates a mitigation solver", + true + ); + + @Override + public String getName() { + return "Mitigate QASM Code Errors"; + } + + @Override + public String getDescription() { + return "Move QASM input to the error mitigators"; + } + + @Override + public List> getSubRoutines() { + return List.of(MITIGATOR_SUBROUTINE); + } + + @Override + public Mono> solve( + String input, + SubRoutineResolver subRoutineResolver, + SolvingProperties properties + ) { + return subRoutineResolver.runSubRoutine(MITIGATOR_SUBROUTINE, input); + } +} diff --git a/src/main/java/edu/kit/provideq/toolbox/circuit/processing/solver/MoveToOptimizationSolver.java b/src/main/java/edu/kit/provideq/toolbox/circuit/processing/solver/MoveToOptimizationSolver.java new file mode 100644 index 00000000..663edbf5 --- /dev/null +++ b/src/main/java/edu/kit/provideq/toolbox/circuit/processing/solver/MoveToOptimizationSolver.java @@ -0,0 +1,45 @@ +package edu.kit.provideq.toolbox.circuit.processing.solver; + +import edu.kit.provideq.toolbox.Solution; +import edu.kit.provideq.toolbox.circuit.processing.results.Result; +import edu.kit.provideq.toolbox.circuit.processing.solver.optimization.OptimizationConfiguration; +import edu.kit.provideq.toolbox.meta.SolvingProperties; +import edu.kit.provideq.toolbox.meta.SubRoutineDefinition; +import edu.kit.provideq.toolbox.meta.SubRoutineResolver; +import java.util.List; +import org.springframework.stereotype.Component; +import reactor.core.publisher.Mono; + +@Component +public class MoveToOptimizationSolver extends CircuitProcessingSolver { + private static final SubRoutineDefinition OPTIMIZER_SUBROUTINE = + new SubRoutineDefinition<>( + OptimizationConfiguration.OPTIMIZATION_CONFIG, + "Creates a optimization solver", + true + ); + + @Override + public String getName() { + return "Optimize QASM Code"; + } + + @Override + public String getDescription() { + return "Move QASM input to the optimizers"; + } + + @Override + public List> getSubRoutines() { + return List.of(OPTIMIZER_SUBROUTINE); + } + + @Override + public Mono> solve( + String input, + SubRoutineResolver subRoutineResolver, + SolvingProperties properties + ) { + return subRoutineResolver.runSubRoutine(OPTIMIZER_SUBROUTINE, input); + } +} diff --git a/src/main/java/edu/kit/provideq/toolbox/circuit/processing/solver/executor/ExecutionSolver.java b/src/main/java/edu/kit/provideq/toolbox/circuit/processing/solver/executor/ExecutionSolver.java new file mode 100644 index 00000000..60761a3f --- /dev/null +++ b/src/main/java/edu/kit/provideq/toolbox/circuit/processing/solver/executor/ExecutionSolver.java @@ -0,0 +1,150 @@ +package edu.kit.provideq.toolbox.circuit.processing.solver.executor; + +import edu.kit.provideq.toolbox.Solution; +import edu.kit.provideq.toolbox.circuit.processing.results.ExecutionResultHelper; +import edu.kit.provideq.toolbox.circuit.processing.results.Result; +import edu.kit.provideq.toolbox.meta.ProblemSolver; +import edu.kit.provideq.toolbox.meta.ProblemType; +import edu.kit.provideq.toolbox.meta.SolvingProperties; +import edu.kit.provideq.toolbox.meta.SubRoutineResolver; +import edu.kit.provideq.toolbox.meta.setting.SolverSetting; +import edu.kit.provideq.toolbox.meta.setting.basic.IntegerSetting; +import edu.kit.provideq.toolbox.meta.setting.basic.SelectSetting; +import edu.kit.provideq.toolbox.process.ProcessRunner; +import edu.kit.provideq.toolbox.process.PythonProcessRunner; +import java.util.List; +import java.util.Optional; +import org.springframework.beans.factory.annotation.Autowired; +import org.springframework.beans.factory.annotation.Value; +import org.springframework.context.ApplicationContext; +import org.springframework.stereotype.Component; +import reactor.core.publisher.Mono; + +@Component +public class ExecutionSolver implements ProblemSolver { + private static final String SETTING_NUMBER_OF_SHOTS = "Number of shots"; + private static final String SETTING_SELECT_SIMULATOR = "Selected Simulator"; + private static final int DEFAULT_NUMBER_OF_SHOTS = 1024; + private static final QuantumSimulator DEFAULT_SIMULATOR = QuantumSimulator.AER; + + private final String scriptPath; + private final String venv; + private final ApplicationContext context; + + @Autowired + public ExecutionSolver( + @Value("${path.circuitprocessing.circuitexecution}") String scriptPath, + @Value("${venv.circuitprocessing.circuitexecution}") String venv, + ApplicationContext context + ) { + this.context = context; + this.venv = venv; + this.scriptPath = scriptPath; + } + + @Override + public String getName() { + return "Execute OpenQASM circuit"; + } + + @Override + public String getDescription() { + return "Execute an OpenQASM circuit"; + } + + @Override + public List getSolverSettings() { + return List.of( + new IntegerSetting( + SETTING_NUMBER_OF_SHOTS, + "The number of shots to run", + 1, + 1000000, + DEFAULT_NUMBER_OF_SHOTS), + new SelectSetting<>( + SETTING_SELECT_SIMULATOR, + "The simulator to run the code with", + List.of(QuantumSimulator.values()), + QuantumSimulator.AER, + QuantumSimulator::getValue + ) + ); + } + + @Override + public Mono> solve( + String input, + SubRoutineResolver subRoutineResolver, + SolvingProperties properties + ) { + var solution = new Solution<>(this); + + int shotNumber = properties.getSetting(SETTING_NUMBER_OF_SHOTS) + .map(IntegerSetting::getValue) + .orElse(DEFAULT_NUMBER_OF_SHOTS); + + QuantumSimulator selectedSimulator = properties + .>getSetting(SETTING_SELECT_SIMULATOR) + .map(s -> s.getSelectedOptionT(QuantumSimulator::fromValue)) + .orElse(DEFAULT_SIMULATOR); + + var processResult = context + .getBean(PythonProcessRunner.class, scriptPath + "executor.py", venv) + .withArguments( + ProcessRunner.INPUT_FILE_PATH, + String.valueOf(shotNumber), + selectedSimulator.getBackendKey() + ) + .writeInputFile(input) + .readOutputString() + .run(getProblemType(), solution.getId()); + + if (processResult.success()) { + solution.complete(); + solution.setSolutionData( + ExecutionResultHelper.createExecutionResult(processResult.output(), Optional.of(input)) + ); + return Mono.just(solution); + } + solution.fail(); + processResult.errorOutput().ifPresent(solution::setDebugData); + return Mono.just(solution); + } + + @Override + public ProblemType getProblemType() { + return ExecutorConfiguration.EXECUTOR_CONFIG; + } + + enum QuantumSimulator { + AER("AerBackend", "aer"), + // PROJECTQ("ProjectQBackend", "projectq"), + QULACS("QulacsBackend", "qulacs"), + AER_NOISY("Aer Noisy Backend (max. 2 qubits)", "aer_noisy"); + + private final String value; + private final String backendKey; + + QuantumSimulator(String value, String backendKey) { + this.value = value; + this.backendKey = backendKey; + } + + public String getValue() { + return value; + } + + public String getBackendKey() { + return backendKey; + } + + public static QuantumSimulator fromValue(String value) { + for (QuantumSimulator simulator : values()) { + if (simulator.value.equals(value)) { + return simulator; + } + } + throw new IllegalArgumentException("Unknown value: " + value); + } + } +} diff --git a/src/main/java/edu/kit/provideq/toolbox/circuit/processing/solver/executor/ExecutorConfiguration.java b/src/main/java/edu/kit/provideq/toolbox/circuit/processing/solver/executor/ExecutorConfiguration.java new file mode 100644 index 00000000..c5d0d670 --- /dev/null +++ b/src/main/java/edu/kit/provideq/toolbox/circuit/processing/solver/executor/ExecutorConfiguration.java @@ -0,0 +1,49 @@ +package edu.kit.provideq.toolbox.circuit.processing.solver.executor; + +import edu.kit.provideq.toolbox.ResourceProvider; +import edu.kit.provideq.toolbox.circuit.processing.results.Result; +import edu.kit.provideq.toolbox.exception.MissingExampleException; +import edu.kit.provideq.toolbox.meta.Problem; +import edu.kit.provideq.toolbox.meta.ProblemManager; +import edu.kit.provideq.toolbox.meta.ProblemType; +import java.io.IOException; +import java.util.Objects; +import java.util.Set; +import org.springframework.context.annotation.Bean; +import org.springframework.context.annotation.Configuration; + +@Configuration +public class ExecutorConfiguration { + public static final ProblemType EXECUTOR_CONFIG = new ProblemType<>( + "circuit-processing-executor", + "A quantum circuit execution problem that runs a given QASM circuit on a quantum backend " + + "and returns the measurement results.", + String.class, + Result.class + ); + + @Bean + ProblemManager getExecutorProblemManager( + ResourceProvider provider, + ExecutionSolver executionSolver + ) { + return new ProblemManager<>( + EXECUTOR_CONFIG, + Set.of(executionSolver), + loadExampleProblems(provider) + ); + } + + private Set> loadExampleProblems(ResourceProvider provider) { + try { + var problemStream = Objects.requireNonNull( + getClass().getResourceAsStream("../../bell-state.qasm"), + "Problem bell-state.qasm not found"); + var problem = new Problem<>(EXECUTOR_CONFIG); + problem.setInput(provider.readStream(problemStream)); + return Set.of(problem); + } catch (IOException e) { + throw new MissingExampleException(EXECUTOR_CONFIG, e); + } + } +} diff --git a/src/main/java/edu/kit/provideq/toolbox/circuit/processing/solver/mitigation/ErrorMitigationConfiguration.java b/src/main/java/edu/kit/provideq/toolbox/circuit/processing/solver/mitigation/ErrorMitigationConfiguration.java new file mode 100644 index 00000000..0e5d5580 --- /dev/null +++ b/src/main/java/edu/kit/provideq/toolbox/circuit/processing/solver/mitigation/ErrorMitigationConfiguration.java @@ -0,0 +1,48 @@ +package edu.kit.provideq.toolbox.circuit.processing.solver.mitigation; + +import edu.kit.provideq.toolbox.ResourceProvider; +import edu.kit.provideq.toolbox.circuit.processing.results.Result; +import edu.kit.provideq.toolbox.exception.MissingExampleException; +import edu.kit.provideq.toolbox.meta.Problem; +import edu.kit.provideq.toolbox.meta.ProblemManager; +import edu.kit.provideq.toolbox.meta.ProblemType; +import java.io.IOException; +import java.util.Objects; +import java.util.Set; +import org.springframework.context.annotation.Bean; +import org.springframework.context.annotation.Configuration; + +@Configuration +public class ErrorMitigationConfiguration { + public static final ProblemType MITIGATION_CONFIG = new ProblemType<>( + "circuit-processing-mitigation", + "A quantum circuit error mitigation problem that applies error mitigation techniques to " + + "a given QASM circuit.", + String.class, + Result.class + ); + + @Bean + ProblemManager getMitigationProblemManager( + ResourceProvider provider, + ErrorMitigationSolver errorMitigationSolver + ) { + return new ProblemManager<>( + MITIGATION_CONFIG, + Set.of(errorMitigationSolver), + loadExampleProblems(provider) + ); + } + + private Set> loadExampleProblems(ResourceProvider provider) { + try { + var problemStream = Objects.requireNonNull( + getClass().getResourceAsStream("../../bell-state.qasm"), "Problem bell-state.qasm not found"); + var problem = new Problem<>(MITIGATION_CONFIG); + problem.setInput(provider.readStream(problemStream)); + return Set.of(problem); + } catch (IOException e) { + throw new MissingExampleException(MITIGATION_CONFIG, e); + } + } +} diff --git a/src/main/java/edu/kit/provideq/toolbox/circuit/processing/solver/mitigation/ErrorMitigationSolver.java b/src/main/java/edu/kit/provideq/toolbox/circuit/processing/solver/mitigation/ErrorMitigationSolver.java new file mode 100644 index 00000000..e4ea8036 --- /dev/null +++ b/src/main/java/edu/kit/provideq/toolbox/circuit/processing/solver/mitigation/ErrorMitigationSolver.java @@ -0,0 +1,42 @@ +package edu.kit.provideq.toolbox.circuit.processing.solver.mitigation; + +import edu.kit.provideq.toolbox.Solution; +import edu.kit.provideq.toolbox.circuit.processing.results.Result; +import edu.kit.provideq.toolbox.circuit.processing.results.StringResult; +import edu.kit.provideq.toolbox.meta.ProblemSolver; +import edu.kit.provideq.toolbox.meta.ProblemType; +import edu.kit.provideq.toolbox.meta.SolvingProperties; +import edu.kit.provideq.toolbox.meta.SubRoutineResolver; +import org.springframework.stereotype.Component; +import reactor.core.publisher.Mono; + +@Component +public class ErrorMitigationSolver implements ProblemSolver { + + @Override + public String getName() { + return "Mitigate Errors for OpenQASM"; + } + + @Override + public String getDescription() { + return "Run error mitigation strategies on an OpenQASM circuit"; + } + + @Override + public Mono> solve( + String input, + SubRoutineResolver subRoutineResolver, + SolvingProperties properties + ) { + var solution = new Solution<>(this); + solution.setSolutionData(new StringResult(input)); + solution.complete(); + return Mono.just(solution); + } + + @Override + public ProblemType getProblemType() { + return ErrorMitigationConfiguration.MITIGATION_CONFIG; + } +} diff --git a/src/main/java/edu/kit/provideq/toolbox/circuit/processing/solver/optimization/OptimizationConfiguration.java b/src/main/java/edu/kit/provideq/toolbox/circuit/processing/solver/optimization/OptimizationConfiguration.java new file mode 100644 index 00000000..16a2f15c --- /dev/null +++ b/src/main/java/edu/kit/provideq/toolbox/circuit/processing/solver/optimization/OptimizationConfiguration.java @@ -0,0 +1,48 @@ +package edu.kit.provideq.toolbox.circuit.processing.solver.optimization; + +import edu.kit.provideq.toolbox.ResourceProvider; +import edu.kit.provideq.toolbox.circuit.processing.results.Result; +import edu.kit.provideq.toolbox.exception.MissingExampleException; +import edu.kit.provideq.toolbox.meta.Problem; +import edu.kit.provideq.toolbox.meta.ProblemManager; +import edu.kit.provideq.toolbox.meta.ProblemType; +import java.io.IOException; +import java.util.Objects; +import java.util.Set; +import org.springframework.context.annotation.Bean; +import org.springframework.context.annotation.Configuration; + +@Configuration +public class OptimizationConfiguration { + public static final ProblemType OPTIMIZATION_CONFIG = new ProblemType<>( + "circuit-processing-optimization", + "A quantum circuit optimization problem that reduces gate count or circuit depth of a " + + "given QASM circuit.", + String.class, + Result.class + ); + + @Bean + ProblemManager getOptimizationProblemManager( + ResourceProvider provider, + OptimizationSolver optimizationSolver + ) { + return new ProblemManager<>( + OPTIMIZATION_CONFIG, + Set.of(optimizationSolver), + loadExampleProblems(provider) + ); + } + + private Set> loadExampleProblems(ResourceProvider provider) { + try { + var problemStream = Objects.requireNonNull( + getClass().getResourceAsStream("../../bell-state.qasm"), "Problem bell-state.qasm not found"); + var problem = new Problem<>(OPTIMIZATION_CONFIG); + problem.setInput(provider.readStream(problemStream)); + return Set.of(problem); + } catch (IOException e) { + throw new MissingExampleException(OPTIMIZATION_CONFIG, e); + } + } +} diff --git a/src/main/java/edu/kit/provideq/toolbox/circuit/processing/solver/optimization/OptimizationSolver.java b/src/main/java/edu/kit/provideq/toolbox/circuit/processing/solver/optimization/OptimizationSolver.java new file mode 100644 index 00000000..c13d727d --- /dev/null +++ b/src/main/java/edu/kit/provideq/toolbox/circuit/processing/solver/optimization/OptimizationSolver.java @@ -0,0 +1,144 @@ +package edu.kit.provideq.toolbox.circuit.processing.solver.optimization; + +import edu.kit.provideq.toolbox.Solution; +import edu.kit.provideq.toolbox.circuit.processing.results.Result; +import edu.kit.provideq.toolbox.circuit.processing.results.StringResult; +import edu.kit.provideq.toolbox.circuit.processing.solver.CircuitProcessingSolver; +import edu.kit.provideq.toolbox.meta.ProblemSolver; +import edu.kit.provideq.toolbox.meta.ProblemType; +import edu.kit.provideq.toolbox.meta.SolvingProperties; +import edu.kit.provideq.toolbox.meta.SubRoutineDefinition; +import edu.kit.provideq.toolbox.meta.SubRoutineResolver; +import edu.kit.provideq.toolbox.meta.setting.SolverSetting; +import edu.kit.provideq.toolbox.meta.setting.basic.SelectSetting; +import edu.kit.provideq.toolbox.process.ProcessRunner; +import edu.kit.provideq.toolbox.process.PythonProcessRunner; +import java.util.List; +import org.springframework.beans.factory.annotation.Autowired; +import org.springframework.beans.factory.annotation.Value; +import org.springframework.context.ApplicationContext; +import org.springframework.stereotype.Component; +import reactor.core.publisher.Mono; + +@Component +public class OptimizationSolver implements ProblemSolver { + private static final String SETTING_SELECT_OPTIMIZER = "Selected Optimization Pass"; + private static final OptimizationSolver.QuantumOptimizer DEFAULT_OPTIMIZER = + QuantumOptimizer.DECOMPOSE_MULTI_CX; + + private final String scriptPath; + private final String venv; + private final ApplicationContext context; + + @Autowired + public OptimizationSolver( + @Value("${path.circuitprocessing.circuitoptimization}") String scriptPath, + @Value("${venv.circuitprocessing.circuitoptimization}") String venv, + ApplicationContext context + ) { + this.scriptPath = scriptPath; + this.venv = venv; + this.context = context; + } + + @Override + public String getName() { + return "Apply Tket Optimization Pass"; + } + + @Override + public String getDescription() { + return "Transform the given circuit into an optimized but equivalent circuit using" + + "Tket compilation passes (e.g. removing redundancies)."; + } + + @Override + public List> getSubRoutines() { + return List.of(CircuitProcessingSolver.CIRCUIT_PROCESSING_SUBROUTINE); + } + + @Override + public List getSolverSettings() { + return List.of( + new SelectSetting<>( + SETTING_SELECT_OPTIMIZER, + "The optimization pass to refactor the code with", + List.of(OptimizationSolver.QuantumOptimizer.values()), + QuantumOptimizer.DECOMPOSE_MULTI_CX, + OptimizationSolver.QuantumOptimizer::getValue + ) + ); + } + + @Override + public Mono> solve( + String input, + SubRoutineResolver subRoutineResolver, + SolvingProperties properties + ) { + var solution = new Solution<>(this); + + OptimizationSolver.QuantumOptimizer selectedOptimizer = properties + .>getSetting(SETTING_SELECT_OPTIMIZER) + .map(s -> s.getSelectedOptionT(OptimizationSolver.QuantumOptimizer::fromValue)) + .orElse(DEFAULT_OPTIMIZER); + + //String[] inputArray = new String[]{input}; + var processResult = context + .getBean(PythonProcessRunner.class, scriptPath + selectedOptimizer.getScriptPath(), venv) + .withArguments( + ProcessRunner.INPUT_FILE_PATH + ) + .writeInputFile(input) + .readOutputString() + .run(getProblemType(), solution.getId()); + + if (processResult.success() && processResult.output().isPresent()) { + solution.complete(); + solution.setSolutionData(new StringResult(processResult.output().get())); + return subRoutineResolver + .runSubRoutine(CircuitProcessingSolver.CIRCUIT_PROCESSING_SUBROUTINE, + processResult.output().get()); + } + solution.fail(); + processResult.errorOutput().ifPresent(solution::setDebugData); + return Mono.just(solution); + } + + @Override + public ProblemType getProblemType() { + return OptimizationConfiguration.OPTIMIZATION_CONFIG; + } + + enum QuantumOptimizer { + DECOMPOSE_MULTI_CX("DecomposeMultiQubitsCX", + "decompose-multi-cx/decompose_multi_cx_optimizer.py"), + REMOVE_REDUNDANCIES("RemoveRedundancies", + "remove-redundancies/remove_redundancies_optimizer.py"); + + private final String value; + private final String scriptPath; + + QuantumOptimizer(String value, String scriptPath) { + this.value = value; + this.scriptPath = scriptPath; + } + + public String getValue() { + return value; + } + + public String getScriptPath() { + return scriptPath; + } + + public static OptimizationSolver.QuantumOptimizer fromValue(String value) { + for (OptimizationSolver.QuantumOptimizer simulator : values()) { + if (simulator.value.equals(value)) { + return simulator; + } + } + throw new IllegalArgumentException("Unknown value: " + value); + } + } +} diff --git a/src/main/java/edu/kit/provideq/toolbox/knapsack/KnapsackConfiguration.java b/src/main/java/edu/kit/provideq/toolbox/knapsack/KnapsackConfiguration.java index e0320cf3..a3f88d9e 100644 --- a/src/main/java/edu/kit/provideq/toolbox/knapsack/KnapsackConfiguration.java +++ b/src/main/java/edu/kit/provideq/toolbox/knapsack/KnapsackConfiguration.java @@ -6,6 +6,7 @@ import edu.kit.provideq.toolbox.exception.MissingExampleException; import edu.kit.provideq.toolbox.knapsack.solvers.PythonKnapsackSolver; import edu.kit.provideq.toolbox.knapsack.solvers.QiskitKnapsackSolver; +import edu.kit.provideq.toolbox.knapsack.solvers.QuantumTreeGeneratorSolver; import edu.kit.provideq.toolbox.meta.Problem; import edu.kit.provideq.toolbox.meta.ProblemManager; import edu.kit.provideq.toolbox.meta.ProblemType; @@ -37,8 +38,8 @@ public class KnapsackConfiguration { for (int i = 1; i < parts.length - 1; i++) { var item = parts[i].split(" "); items.add(new AbstractMap.SimpleEntry<>( - Integer.parseInt(item[1]), - Integer.parseInt(item[2])) + Integer.parseInt(item[1]), + Integer.parseInt(item[2])) ); } items.sort(Comparator.comparingInt(a -> -a.getKey() / a.getValue())); @@ -88,14 +89,15 @@ public class KnapsackConfiguration { @Bean ProblemManager getKnapsackManager( - PythonKnapsackSolver pythonKnapsackSolver, - QiskitKnapsackSolver qiskitKnapsackSolver, - ResourceProvider resourceProvider + PythonKnapsackSolver pythonKnapsackSolver, + QiskitKnapsackSolver qiskitKnapsackSolver, + QuantumTreeGeneratorSolver quantumTreeGeneratorSolver, + ResourceProvider resourceProvider ) { return new ProblemManager<>( - KNAPSACK, - Set.of(pythonKnapsackSolver, qiskitKnapsackSolver), - loadExampleProblems(resourceProvider) + KNAPSACK, + Set.of(pythonKnapsackSolver, qiskitKnapsackSolver, quantumTreeGeneratorSolver), + loadExampleProblems(resourceProvider) ); } diff --git a/src/main/java/edu/kit/provideq/toolbox/knapsack/solvers/QuantumTreeGeneratorSolver.java b/src/main/java/edu/kit/provideq/toolbox/knapsack/solvers/QuantumTreeGeneratorSolver.java new file mode 100644 index 00000000..9dc73339 --- /dev/null +++ b/src/main/java/edu/kit/provideq/toolbox/knapsack/solvers/QuantumTreeGeneratorSolver.java @@ -0,0 +1,93 @@ +package edu.kit.provideq.toolbox.knapsack.solvers; + +import edu.kit.provideq.toolbox.Solution; +import edu.kit.provideq.toolbox.SolutionStatus; +import edu.kit.provideq.toolbox.circuit.processing.CircuitProcessingConfiguration; +import edu.kit.provideq.toolbox.circuit.processing.results.ExecutionResultVisitor; +import edu.kit.provideq.toolbox.circuit.processing.results.Result; +import edu.kit.provideq.toolbox.meta.SolvingProperties; +import edu.kit.provideq.toolbox.meta.SubRoutineDefinition; +import edu.kit.provideq.toolbox.meta.SubRoutineResolver; +import edu.kit.provideq.toolbox.process.ProcessRunner; +import edu.kit.provideq.toolbox.process.PythonProcessRunner; +import java.util.List; +import org.springframework.beans.factory.annotation.Autowired; +import org.springframework.beans.factory.annotation.Value; +import org.springframework.context.ApplicationContext; +import org.springframework.stereotype.Component; +import reactor.core.publisher.Mono; + +@Component +public class QuantumTreeGeneratorSolver extends KnapsackSolver { + private static final SubRoutineDefinition CIRCUIT_PROCESSING_SUBROUTINE = + new SubRoutineDefinition<>( + CircuitProcessingConfiguration.CIRCUIT_PROCESSING, + "Use circuit processing", + true + ); + + + private final String scriptPath; + private final String venv; + private final ApplicationContext context; + + @Autowired + public QuantumTreeGeneratorSolver( + @Value("${path.qiskit.knapsack_quantum_tree_generator}") String scriptPath, + @Value("${venv.qiskit.knapsack_quantum_tree_generator}") String venv, + ApplicationContext context) { + this.scriptPath = scriptPath; + this.venv = venv; + this.context = context; + } + + @Override + public String getName() { + return "Knapsack Quantum Tree Generator"; + } + + @Override + public String getDescription() { + return "Solve Knapsack using the Quantum Tree Generator solver."; + } + + @Override + public Mono> solve( + String input, + SubRoutineResolver subRoutineResolver, + SolvingProperties properties + ) { + var solution = new Solution<>(this); + + var processResult = context + .getBean(PythonProcessRunner.class, scriptPath, venv) + .withArguments( + ProcessRunner.INPUT_FILE_PATH, + ProcessRunner.OUTPUT_FILE_PATH + ) + .writeInputFile(input) + .readOutputFile() + .run(getProblemType(), solution.getId()); + + String openQasm = processResult.output().orElseThrow(); + return subRoutineResolver.runSubRoutine(CIRCUIT_PROCESSING_SUBROUTINE, openQasm) + .map(resultSolution -> { + Solution s = new Solution<>(this); + SolutionStatus status = resultSolution.getStatus(); + if (status == SolutionStatus.ERROR) { + s.fail(); + s.setDebugData(resultSolution.getDebugData()); + return s; + } + + s.complete(); + s.setSolutionData(resultSolution.getSolutionData().accept(new ExecutionResultVisitor())); + return s; + }); + } + + @Override + public List> getSubRoutines() { + return List.of(CIRCUIT_PROCESSING_SUBROUTINE); + } +} diff --git a/src/main/java/edu/kit/provideq/toolbox/tools/equivalencechecking/EquivalenceChecking.java b/src/main/java/edu/kit/provideq/toolbox/tools/equivalencechecking/EquivalenceChecking.java new file mode 100644 index 00000000..187ac1c9 --- /dev/null +++ b/src/main/java/edu/kit/provideq/toolbox/tools/equivalencechecking/EquivalenceChecking.java @@ -0,0 +1,81 @@ +package edu.kit.provideq.toolbox.tools.equivalencechecking; + +import com.fasterxml.jackson.core.JsonProcessingException; +import com.fasterxml.jackson.databind.JsonNode; +import com.fasterxml.jackson.databind.ObjectMapper; +import edu.kit.provideq.toolbox.meta.ProblemType; +import edu.kit.provideq.toolbox.process.ProcessRunner; +import edu.kit.provideq.toolbox.process.PythonProcessRunner; +import java.util.UUID; +import org.springframework.beans.factory.annotation.Autowired; +import org.springframework.beans.factory.annotation.Value; +import org.springframework.context.ApplicationContext; +import org.springframework.stereotype.Component; + +/** Runs the circuit equivalence-checking Python tool. */ +@Component +public class EquivalenceChecking { + private static final ProblemType TOOL_TYPE = new ProblemType<>( + "equivalencechecking", + "Circuit equivalence checking", + String.class, + String.class + ); + + private final String scriptPath; + private final String venv; + private final ApplicationContext context; + private final ObjectMapper objectMapper; + + /** Creates an equivalence-checking tool backed by the configured Python script. */ + @Autowired + public EquivalenceChecking( + @Value("${path.tools.equivalencechecking}") String scriptPath, + @Value("${venv.tools.equivalencechecking}") String venv, + ApplicationContext context, + ObjectMapper objectMapper) { + this.scriptPath = scriptPath; + this.venv = venv; + this.context = context; + this.objectMapper = objectMapper; + } + + /** Executes the Python tool and returns its JSON response. */ + public JsonNode check(JsonNode input) { + final String serializedInput; + try { + serializedInput = objectMapper.writeValueAsString(input); + } catch (JsonProcessingException exception) { + throw new IllegalArgumentException("Could not serialize the request JSON.", exception); + } + + var processResult = context + .getBean(PythonProcessRunner.class, scriptPath, venv) + .withArguments( + ProcessRunner.INPUT_FILE_PATH, + ProcessRunner.OUTPUT_FILE_PATH + ) + .writeInputFile(serializedInput) + .readOutputFile() + .run(TOOL_TYPE, UUID.randomUUID()); + + if (!processResult.success()) { + throw new IllegalStateException(processResult.errorOutput() + .orElse("Equivalence checking failed without an error message.")); + } + + var output = processResult.output() + .orElseThrow(() -> new IllegalStateException( + "Equivalence checking completed without returning output.")); + try { + var jsonOutput = objectMapper.readTree(output); + if (jsonOutput == null) { + throw new IllegalStateException("Equivalence checking returned empty output."); + } + return jsonOutput; + } catch (JsonProcessingException exception) { + throw new IllegalStateException( + "Equivalence checking returned invalid JSON.", exception); + } + } +} diff --git a/src/main/java/edu/kit/provideq/toolbox/util/Pair.java b/src/main/java/edu/kit/provideq/toolbox/util/Pair.java new file mode 100644 index 00000000..7c523ca5 --- /dev/null +++ b/src/main/java/edu/kit/provideq/toolbox/util/Pair.java @@ -0,0 +1,17 @@ +package edu.kit.provideq.toolbox.util; + +/** + * Represents an immutable pair of two values. + * + * @param The type of the first element in the pair. + * @param The type of the second element in the pair. + */ +public record Pair(T1 first, T2 second) { + public Pair copyWithFirst(T1 first) { + return new Pair<>(first, this.second); + } + + public Pair copyWithSecond(T2 second) { + return new Pair<>(this.first, second); + } +} diff --git a/src/main/resources/application.properties b/src/main/resources/application.properties index 01cdfc62..9c534643 100644 --- a/src/main/resources/application.properties +++ b/src/main/resources/application.properties @@ -1 +1 @@ -# default spring profile, correct one will be set during runtime (see ToolboxServerApplication.java) # options: mac, windows, linux spring.profiles.active=linux springdoc.swagger-ui.operationsSorter=alpha springdoc.swagger-ui.tagsSorter=alpha working.directory=jobs examples.directory=examples springdoc.swagger-ui.path=/ # Solvers name.solvers=solvers # Non OS-specific solvers: (typically GAMS and Python) name.gams=gams path.gams=${name.solvers}/${name.gams} name.gams.max-cut=max-cut path.gams.max-cut=${path.gams}/${name.gams.max-cut}/maxcut.gms name.gams.sat=sat path.gams.sat=${path.gams}/${name.gams.sat}/sat.gms name.qiskit=qiskit path.qiskit=${name.solvers}/${name.qiskit} name.qiskit.knapsack=knapsack path.qiskit.knapsack=${path.qiskit}/${name.qiskit.knapsack}/knapsack_qiskit.py venv.qiskit.knapsack=${name.solvers}_${name.qiskit}_${name.qiskit.knapsack} name.qiskit.materialsimulation=materialsimulation path.qiskit.materialsimulation=${path.qiskit}/${name.qiskit.materialsimulation}/material_simulation_qiskit.py venv.qiskit.materialsimulation=${name.solvers}_${name.qiskit}_${name.qiskit.materialsimulation} name.qiskit.max-cut=max-cut path.qiskit.max-cut=${path.qiskit}/${name.qiskit.max-cut}/maxCut_qiskit.py venv.qiskit.max-cut=${name.solvers}_${name.qiskit}_${name.qiskit.max-cut} name.qiskit.qubo=qubo path.qiskit.qubo=${path.qiskit}/${name.qiskit.qubo}/qubo_qiskit.py venv.qiskit.qubo=${name.solvers}_${name.qiskit}_${name.qiskit.qubo} name.cirq=cirq path.cirq=${name.solvers}/${name.cirq} name.cirq.max-cut=max-cut path.cirq.max-cut=${path.cirq}/${name.cirq.max-cut}/max_cut_cirq.py venv.cirq.max-cut=${name.solvers}_${name.cirq}_${name.cirq.max-cut} name.qrisp=qrisp path.qrisp=${name.solvers}/${name.qrisp} name.qrisp.vrp=vrp path.qrisp.vrp=${path.qrisp}/${name.qrisp.vrp}/grover.py venv.qrisp.vrp=${name.solvers}_${name.qrisp}_${name.qrisp.vrp} name.qrisp.qubo=qubo path.qrisp.qubo=${path.qrisp}/${name.qrisp.qubo}/qaoa.py venv.qrisp.qubo=${name.solvers}_${name.qrisp}_${name.qrisp.qubo} name.qrisp.sat=sat path.qrisp.sat.grover=${path.qrisp}/${name.qrisp.sat}/grover.py path.qrisp.sat.exact=${path.qrisp}/${name.qrisp.sat}/exact_grover.py venv.qrisp.sat=${name.solvers}_${name.qrisp}_${name.qrisp.sat} name.dwave=dwave path.dwave=${name.solvers}/${name.dwave} name.dwave.qubo=qubo path.dwave.qubo=${path.dwave}/${name.dwave.qubo}/main.py venv.dwave.qubo=${name.solvers}_${name.dwave}_${name.dwave.qubo} # Non OS-specific custom solvers: (solvers that are not part of a framework) name.custom=custom path.custom=${name.solvers}/${name.custom} name.custom.hs-knapsack=hs-knapsack path.custom.hs-knapsack=${path.custom}/${name.custom.hs-knapsack}/knapsack.py venv.custom.hs-knapsack=${name.solvers}_${name.custom}_${name.custom.hs-knapsack} name.custom.lkh=lkh path.custom.lkh=${path.custom}/${name.custom.lkh}/vrp_lkh.py venv.custom.lkh=${name.solvers}_${name.custom}_${name.custom.lkh} name.custom.berger-vrp=berger-vrp name.custom.sharp-sat-bruteforce=sharp-sat-bruteforce path.custom.sharp-sat-bruteforce=${path.custom}/${name.custom.sharp-sat-bruteforce}/exact-solution-counter.py venv.custom.sharp-sat-bruteforce=${name.solvers}_${name.custom}_${name.custom.sharp-sat-bruteforce} name.custom.sharp-sat-ganak=sharp-sat-ganak venv.custom.sharp-sat-ganak=${name.solvers}_${name.custom}_${name.custom.sharp-sat-ganak} # Demonstrators name.demonstrators=demonstrators name.demonstrators.cplex=cplex path.demonstrators.cplex=${name.demonstrators}/${name.demonstrators.cplex} name.demonstrators.cplex.mip=mip-solver path.demonstrators.cplex.mip=${path.demonstrators.cplex}/${name.demonstrators.cplex.mip}/mip-solver.py venv.demonstrators.cplex.mip=${name.demonstrators}_${name.demonstrators.cplex}_${name.demonstrators.cplex.mip} name.demonstrators.qiskit=qiskit path.demonstrators.qiskit=${name.demonstrators}/${name.demonstrators.qiskit} name.demonstrators.qiskit.molecule-energy=molecule-energy path.demonstrators.qiskit.molecule-energy=${path.demonstrators.qiskit}/${name.demonstrators.qiskit.molecule-energy}/molecule-energy.py venv.demonstrators.qiskit.molecule-energy=${name.demonstrators}_${name.demonstrators.qiskit}_${name.demonstrators.qiskit.molecule-energy} \ No newline at end of file +# default spring profile, correct one will be set during runtime (see ToolboxServerApplication.java) # options: mac, windows, linux spring.profiles.active=linux springdoc.swagger-ui.operationsSorter=alpha springdoc.swagger-ui.tagsSorter=alpha working.directory=jobs examples.directory=examples springdoc.swagger-ui.path=/ # Solvers name.solvers=solvers # Non OS-specific solvers: (typically GAMS and Python) name.gams=gams path.gams=${name.solvers}/${name.gams} name.gams.max-cut=max-cut path.gams.max-cut=${path.gams}/${name.gams.max-cut}/maxcut.gms name.gams.sat=sat path.gams.sat=${path.gams}/${name.gams.sat}/sat.gms name.qiskit=qiskit path.qiskit=${name.solvers}/${name.qiskit} name.qiskit.knapsack=knapsack path.qiskit.knapsack=${path.qiskit}/${name.qiskit.knapsack}/knapsack_qiskit.py venv.qiskit.knapsack=${name.solvers}_${name.qiskit}_${name.qiskit.knapsack} name.qiskit.knapsack_quantum_tree_generator=knapsack_quantum_tree_generator path.qiskit.knapsack_quantum_tree_generator=${path.qiskit}/${name.qiskit.knapsack_quantum_tree_generator}/knapsack_quantum_tree_generator_openqasm.py venv.qiskit.knapsack_quantum_tree_generator=${name.solvers}_${name.qiskit}_${name.qiskit.knapsack_quantum_tree_generator} name.qiskit.materialsimulation=materialsimulation path.qiskit.materialsimulation=${path.qiskit}/${name.qiskit.materialsimulation}/material_simulation_qiskit.py venv.qiskit.materialsimulation=${name.solvers}_${name.qiskit}_${name.qiskit.materialsimulation} name.qiskit.max-cut=max-cut path.qiskit.max-cut=${path.qiskit}/${name.qiskit.max-cut}/maxCut_qiskit.py venv.qiskit.max-cut=${name.solvers}_${name.qiskit}_${name.qiskit.max-cut} name.qiskit.qubo=qubo path.qiskit.qubo=${path.qiskit}/${name.qiskit.qubo}/qubo_qiskit.py venv.qiskit.qubo=${name.solvers}_${name.qiskit}_${name.qiskit.qubo} name.cirq=cirq path.cirq=${name.solvers}/${name.cirq} name.cirq.max-cut=max-cut path.cirq.max-cut=${path.cirq}/${name.cirq.max-cut}/max_cut_cirq.py venv.cirq.max-cut=${name.solvers}_${name.cirq}_${name.cirq.max-cut} name.qrisp=qrisp path.qrisp=${name.solvers}/${name.qrisp} name.qrisp.vrp=vrp path.qrisp.vrp=${path.qrisp}/${name.qrisp.vrp}/grover.py venv.qrisp.vrp=${name.solvers}_${name.qrisp}_${name.qrisp.vrp} name.qrisp.qubo=qubo path.qrisp.qubo=${path.qrisp}/${name.qrisp.qubo}/qaoa.py venv.qrisp.qubo=${name.solvers}_${name.qrisp}_${name.qrisp.qubo} name.qrisp.sat=sat path.qrisp.sat.grover=${path.qrisp}/${name.qrisp.sat}/grover.py path.qrisp.sat.exact=${path.qrisp}/${name.qrisp.sat}/exact_grover.py venv.qrisp.sat=${name.solvers}_${name.qrisp}_${name.qrisp.sat} name.dwave=dwave path.dwave=${name.solvers}/${name.dwave} name.dwave.qubo=qubo path.dwave.qubo=${path.dwave}/${name.dwave.qubo}/main.py venv.dwave.qubo=${name.solvers}_${name.dwave}_${name.dwave.qubo} name.tools=tools path.tools=${name.solvers}/${name.tools} name.tools.equivalencechecking=equivalencechecking path.tools.equivalencechecking=${path.tools}/${name.tools.equivalencechecking}/equivalencechecking.py venv.tools.equivalencechecking=${name.solvers}_${name.tools}_${name.tools.equivalencechecking} # Non OS-specific custom solvers: (solvers that are not part of a framework) name.custom=custom path.custom=${name.solvers}/${name.custom} name.custom.hs-knapsack=hs-knapsack path.custom.hs-knapsack=${path.custom}/${name.custom.hs-knapsack}/knapsack.py venv.custom.hs-knapsack=${name.solvers}_${name.custom}_${name.custom.hs-knapsack} name.custom.lkh=lkh path.custom.lkh=${path.custom}/${name.custom.lkh}/vrp_lkh.py venv.custom.lkh=${name.solvers}_${name.custom}_${name.custom.lkh} name.custom.berger-vrp=berger-vrp name.custom.sharp-sat-bruteforce=sharp-sat-bruteforce path.custom.sharp-sat-bruteforce=${path.custom}/${name.custom.sharp-sat-bruteforce}/exact-solution-counter.py venv.custom.sharp-sat-bruteforce=${name.solvers}_${name.custom}_${name.custom.sharp-sat-bruteforce} name.custom.sharp-sat-ganak=sharp-sat-ganak venv.custom.sharp-sat-ganak=${name.solvers}_${name.custom}_${name.custom.sharp-sat-ganak} # Demonstrators name.demonstrators=demonstrators name.demonstrators.cplex=cplex path.demonstrators.cplex=${name.demonstrators}/${name.demonstrators.cplex} name.demonstrators.cplex.mip=mip-solver path.demonstrators.cplex.mip=${path.demonstrators.cplex}/${name.demonstrators.cplex.mip}/mip-solver.py venv.demonstrators.cplex.mip=${name.demonstrators}_${name.demonstrators.cplex}_${name.demonstrators.cplex.mip} name.demonstrators.qiskit=qiskit path.demonstrators.qiskit=${name.demonstrators}/${name.demonstrators.qiskit} name.demonstrators.qiskit.molecule-energy=molecule-energy path.demonstrators.qiskit.molecule-energy=${path.demonstrators.qiskit}/${name.demonstrators.qiskit.molecule-energy}/molecule-energy.py venv.demonstrators.qiskit.molecule-energy=${name.demonstrators}_${name.demonstrators.qiskit}_${name.demonstrators.qiskit.molecule-energy} name.circuitprocessing=circuitprocessing path.circuitprocessing=${name.solvers}/${name.circuitprocessing} name.circuitprocessing.circuitexecution=circuitexecution path.circuitprocessing.circuitexecution=${path.circuitprocessing}/${name.circuitprocessing.circuitexecution}/ venv.circuitprocessing.circuitexecution=${name.solvers}_${name.circuitprocessing}_${name.circuitprocessing.circuitexecution} name.circuitprocessing.circuitoptimization=circuitoptimizing path.circuitprocessing.circuitoptimization=${path.circuitprocessing}/${name.circuitprocessing.circuitoptimization}/ venv.circuitprocessing.circuitoptimization=${name.solvers}_${name.circuitprocessing}_${name.circuitprocessing.circuitoptimization} \ No newline at end of file diff --git a/src/main/resources/edu/kit/provideq/toolbox/circuit/processing/bell-state.qasm b/src/main/resources/edu/kit/provideq/toolbox/circuit/processing/bell-state.qasm new file mode 100644 index 00000000..c3bf3c2b --- /dev/null +++ b/src/main/resources/edu/kit/provideq/toolbox/circuit/processing/bell-state.qasm @@ -0,0 +1,8 @@ +OPENQASM 2.0; +include "qelib1.inc"; +qreg q[2]; +creg c[2]; +h q[0]; +cx q[0],q[1]; +measure q[0] -> c[0]; +measure q[1] -> c[1]; diff --git a/src/main/resources/edu/kit/provideq/toolbox/circuit/processing/cswap.qasm b/src/main/resources/edu/kit/provideq/toolbox/circuit/processing/cswap.qasm new file mode 100644 index 00000000..b2bce4dc --- /dev/null +++ b/src/main/resources/edu/kit/provideq/toolbox/circuit/processing/cswap.qasm @@ -0,0 +1,6 @@ +OPENQASM 2.0; +include "qelib1.inc"; +qreg q[3]; +crz(0.5) q[0], q[1]; +t q[2]; +cswap q[2], q[0], q[1]; diff --git a/src/test/java/edu/kit/provideq/toolbox/api/CircuitProcessingSolversTest.java b/src/test/java/edu/kit/provideq/toolbox/api/CircuitProcessingSolversTest.java new file mode 100644 index 00000000..e31e8382 --- /dev/null +++ b/src/test/java/edu/kit/provideq/toolbox/api/CircuitProcessingSolversTest.java @@ -0,0 +1,112 @@ +package edu.kit.provideq.toolbox.api; + +import static edu.kit.provideq.toolbox.circuit.processing.CircuitProcessingConfiguration.CIRCUIT_PROCESSING; +import static edu.kit.provideq.toolbox.circuit.processing.solver.executor.ExecutorConfiguration.EXECUTOR_CONFIG; +import static edu.kit.provideq.toolbox.circuit.processing.solver.mitigation.ErrorMitigationConfiguration.MITIGATION_CONFIG; +import static edu.kit.provideq.toolbox.circuit.processing.solver.optimization.OptimizationConfiguration.OPTIMIZATION_CONFIG; +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertTrue; + +import edu.kit.provideq.toolbox.circuit.processing.solver.MoveToExecutionSolver; +import edu.kit.provideq.toolbox.circuit.processing.solver.MoveToMitigationSolver; +import edu.kit.provideq.toolbox.circuit.processing.solver.MoveToOptimizationSolver; +import edu.kit.provideq.toolbox.circuit.processing.solver.executor.ExecutionSolver; +import edu.kit.provideq.toolbox.circuit.processing.solver.mitigation.ErrorMitigationSolver; +import edu.kit.provideq.toolbox.circuit.processing.solver.optimization.OptimizationSolver; +import edu.kit.provideq.toolbox.meta.Problem; +import edu.kit.provideq.toolbox.meta.ProblemManager; +import edu.kit.provideq.toolbox.meta.ProblemManagerProvider; +import java.time.Duration; +import java.util.List; +import java.util.Optional; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.springframework.beans.factory.annotation.Autowired; +import org.springframework.boot.test.autoconfigure.web.servlet.AutoConfigureMockMvc; +import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.test.web.reactive.server.WebTestClient; + +@SpringBootTest +@AutoConfigureMockMvc +class CircuitProcessingSolversTest { + @Autowired + private WebTestClient client; + + @Autowired + private ProblemManagerProvider problemManagerProvider; + + @Autowired + private MoveToMitigationSolver moveToMitigationSolver; + + @Autowired + private MoveToExecutionSolver moveToExecutionSolver; + + @Autowired + private MoveToOptimizationSolver moveToOptimizationSolver; + + @Autowired + private ErrorMitigationSolver errorMitigationSolver; + + @Autowired + private ExecutionSolver executionSolver; + + @Autowired + private OptimizationSolver optimizationSolver; + + private ProblemManager problemManager; + private List problems; + + @BeforeEach + void beforeEach() { + this.client = this.client.mutate() + .responseTimeout(Duration.ofSeconds(60)) + .build(); + problemManager = problemManagerProvider.findProblemManagerForType(CIRCUIT_PROCESSING).get(); + problems = problemManager.getExampleInstances() + .stream() + .map(Problem::getInput) + .filter(Optional::isPresent) + .map(Optional::get) + .toList(); + } + + @Test + void testMoveToMitigationSolver() { + var circuit = problems.get(0); + var problemDto = ApiTestHelper.createProblem(client, moveToMitigationSolver, circuit, CIRCUIT_PROCESSING); + var subProblemId = problemDto.getSubProblems().get(0).getSubProblemIds().get(0); + ApiTestHelper.setProblemSolver(client, errorMitigationSolver, subProblemId, MITIGATION_CONFIG.getId()); + var solvedDto = ApiTestHelper.trySolveFor(60, client, problemDto.getId(), CIRCUIT_PROCESSING); + ApiTestHelper.testSolution(solvedDto); + assertEquals(circuit, solvedDto.getSolution().getSolutionData()); + } + + @Test + void testMoveToExecutionSolver() { + var circuit = problems.get(0); + var problemDto = ApiTestHelper.createProblem(client, moveToExecutionSolver, circuit, CIRCUIT_PROCESSING); + var subProblemId = problemDto.getSubProblems().get(0).getSubProblemIds().get(0); + ApiTestHelper.setProblemSolver(client, executionSolver, subProblemId, EXECUTOR_CONFIG.getId()); + var solvedDto = ApiTestHelper.trySolveFor(60, client, problemDto.getId(), CIRCUIT_PROCESSING); + ApiTestHelper.testSolution(solvedDto); + assertFalse(solvedDto.getSolution().getSolutionData().isBlank()); + } + + @Test + void testMoveToOptimizationSolver() { + var circuit = problems.get(0); + var problemDto = ApiTestHelper.createProblem(client, moveToOptimizationSolver, circuit, CIRCUIT_PROCESSING); + var optSubProblemId = problemDto.getSubProblems().get(0).getSubProblemIds().get(0); + var optDto = ApiTestHelper.setProblemSolver( + client, optimizationSolver, optSubProblemId, OPTIMIZATION_CONFIG.getId()); + var circuitSubProblemId = optDto.getSubProblems().get(0).getSubProblemIds().get(0); + var mitigationEntryDto = ApiTestHelper.setProblemSolver( + client, moveToMitigationSolver, circuitSubProblemId, CIRCUIT_PROCESSING.getId()); + var mitigationId = mitigationEntryDto.getSubProblems().get(0).getSubProblemIds().get(0); + ApiTestHelper.setProblemSolver(client, errorMitigationSolver, mitigationId, MITIGATION_CONFIG.getId()); + var solvedDto = ApiTestHelper.trySolveFor(120, client, problemDto.getId(), CIRCUIT_PROCESSING); + ApiTestHelper.testSolution(solvedDto); + assertTrue(solvedDto.getSolution().getSolutionData().contains("OPENQASM")); + } +} diff --git a/src/test/java/edu/kit/provideq/toolbox/api/ErrorMitigationSolverTest.java b/src/test/java/edu/kit/provideq/toolbox/api/ErrorMitigationSolverTest.java new file mode 100644 index 00000000..1815a325 --- /dev/null +++ b/src/test/java/edu/kit/provideq/toolbox/api/ErrorMitigationSolverTest.java @@ -0,0 +1,52 @@ +package edu.kit.provideq.toolbox.api; + +import static edu.kit.provideq.toolbox.circuit.processing.solver.mitigation.ErrorMitigationConfiguration.MITIGATION_CONFIG; +import static org.junit.jupiter.api.Assertions.assertEquals; + +import edu.kit.provideq.toolbox.circuit.processing.solver.mitigation.ErrorMitigationSolver; +import edu.kit.provideq.toolbox.meta.Problem; +import edu.kit.provideq.toolbox.meta.ProblemManager; +import edu.kit.provideq.toolbox.meta.ProblemManagerProvider; +import java.util.List; +import java.util.Optional; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.springframework.beans.factory.annotation.Autowired; +import org.springframework.boot.test.autoconfigure.web.servlet.AutoConfigureMockMvc; +import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.test.web.reactive.server.WebTestClient; + +@SpringBootTest +@AutoConfigureMockMvc +class ErrorMitigationSolverTest { + @Autowired + private WebTestClient client; + + @Autowired + private ProblemManagerProvider problemManagerProvider; + + @Autowired + private ErrorMitigationSolver errorMitigationSolver; + + private ProblemManager problemManager; + private List problems; + + @BeforeEach + void beforeEach() { + problemManager = problemManagerProvider.findProblemManagerForType(MITIGATION_CONFIG).get(); + problems = problemManager.getExampleInstances() + .stream() + .map(Problem::getInput) + .filter(Optional::isPresent) + .map(Optional::get) + .toList(); + } + + @Test + void testErrorMitigationSolver() { + var circuit = problems.get(0); + var problem = ApiTestHelper.createProblem(client, errorMitigationSolver, circuit, MITIGATION_CONFIG); + ApiTestHelper.testSolution(problem); + assertEquals(circuit, problem.getSolution().getSolutionData()); + } +} diff --git a/src/test/java/edu/kit/provideq/toolbox/api/ExecutionSolverTest.java b/src/test/java/edu/kit/provideq/toolbox/api/ExecutionSolverTest.java new file mode 100644 index 00000000..11549377 --- /dev/null +++ b/src/test/java/edu/kit/provideq/toolbox/api/ExecutionSolverTest.java @@ -0,0 +1,56 @@ +package edu.kit.provideq.toolbox.api; + +import static edu.kit.provideq.toolbox.circuit.processing.solver.executor.ExecutorConfiguration.EXECUTOR_CONFIG; +import static org.junit.jupiter.api.Assertions.assertTrue; + +import edu.kit.provideq.toolbox.Solution; +import edu.kit.provideq.toolbox.circuit.processing.solver.executor.ExecutionSolver; +import edu.kit.provideq.toolbox.meta.Problem; +import edu.kit.provideq.toolbox.meta.ProblemManagerProvider; +import java.time.Duration; +import java.util.List; +import java.util.Optional; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.springframework.beans.factory.annotation.Autowired; +import org.springframework.boot.test.autoconfigure.web.servlet.AutoConfigureMockMvc; +import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.test.web.reactive.server.WebTestClient; + +@SpringBootTest +@AutoConfigureMockMvc +class ExecutionSolverTest { + @Autowired + private WebTestClient client; + + @Autowired + private ProblemManagerProvider problemManagerProvider; + + @Autowired + private ExecutionSolver executionSolver; + + private List problems; + + @BeforeEach + void beforeEach() { + this.client = this.client.mutate() + .responseTimeout(Duration.ofSeconds(60)) + .build(); + problems = problemManagerProvider.findProblemManagerForType(EXECUTOR_CONFIG).get() + .getExampleInstances() + .stream() + .map(Problem::getInput) + .filter(Optional::isPresent) + .map(Optional::get) + .toList(); + } + + @Test + void testExecutionSolver() { + var circuit = problems.get(0); + var problem = ApiTestHelper.createProblem(client, executionSolver, circuit, EXECUTOR_CONFIG); + ApiTestHelper.testSolution(problem); + Solution solution = problem.getSolution(); + assertTrue(solution.getSolutionData().toString().contains("OPENQASM")); + } +} diff --git a/src/test/java/edu/kit/provideq/toolbox/api/OptimizationSolverTest.java b/src/test/java/edu/kit/provideq/toolbox/api/OptimizationSolverTest.java new file mode 100644 index 00000000..4945a5c7 --- /dev/null +++ b/src/test/java/edu/kit/provideq/toolbox/api/OptimizationSolverTest.java @@ -0,0 +1,72 @@ +package edu.kit.provideq.toolbox.api; + +import static edu.kit.provideq.toolbox.circuit.processing.CircuitProcessingConfiguration.CIRCUIT_PROCESSING; +import static edu.kit.provideq.toolbox.circuit.processing.solver.mitigation.ErrorMitigationConfiguration.MITIGATION_CONFIG; +import static edu.kit.provideq.toolbox.circuit.processing.solver.optimization.OptimizationConfiguration.OPTIMIZATION_CONFIG; +import static org.junit.jupiter.api.Assertions.assertTrue; + +import edu.kit.provideq.toolbox.circuit.processing.solver.MoveToMitigationSolver; +import edu.kit.provideq.toolbox.circuit.processing.solver.mitigation.ErrorMitigationSolver; +import edu.kit.provideq.toolbox.circuit.processing.solver.optimization.OptimizationSolver; +import edu.kit.provideq.toolbox.meta.Problem; +import edu.kit.provideq.toolbox.meta.ProblemManager; +import edu.kit.provideq.toolbox.meta.ProblemManagerProvider; +import java.time.Duration; +import java.util.List; +import java.util.Optional; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.springframework.beans.factory.annotation.Autowired; +import org.springframework.boot.test.autoconfigure.web.servlet.AutoConfigureMockMvc; +import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.test.web.reactive.server.WebTestClient; + +@SpringBootTest +@AutoConfigureMockMvc +class OptimizationSolverTest { + @Autowired + private WebTestClient client; + + @Autowired + private ProblemManagerProvider problemManagerProvider; + + @Autowired + private OptimizationSolver optimizationSolver; + + @Autowired + private MoveToMitigationSolver moveToMitigationSolver; + + @Autowired + private ErrorMitigationSolver errorMitigationSolver; + + private ProblemManager problemManager; + private List problems; + + @BeforeEach + void beforeEach() { + this.client = this.client.mutate() + .responseTimeout(Duration.ofSeconds(60)) + .build(); + problemManager = problemManagerProvider.findProblemManagerForType(OPTIMIZATION_CONFIG).get(); + problems = problemManager.getExampleInstances() + .stream() + .map(Problem::getInput) + .filter(Optional::isPresent) + .map(Optional::get) + .toList(); + } + + @Test + void testOptimizationSolver() { + var circuit = problems.get(0); + var problemDto = ApiTestHelper.createProblem(client, optimizationSolver, circuit, OPTIMIZATION_CONFIG); + var circuitSubProblemId = problemDto.getSubProblems().get(0).getSubProblemIds().get(0); + var mitigationEntryDto = ApiTestHelper.setProblemSolver( + client, moveToMitigationSolver, circuitSubProblemId, CIRCUIT_PROCESSING.getId()); + var mitigationId = mitigationEntryDto.getSubProblems().get(0).getSubProblemIds().get(0); + ApiTestHelper.setProblemSolver(client, errorMitigationSolver, mitigationId, MITIGATION_CONFIG.getId()); + var solvedDto = ApiTestHelper.trySolveFor(120, client, problemDto.getId(), OPTIMIZATION_CONFIG); + ApiTestHelper.testSolution(solvedDto); + assertTrue(solvedDto.getSolution().getSolutionData().contains("OPENQASM")); + } +} diff --git a/src/test/java/edu/kit/provideq/toolbox/api/tools/EquivalenceCheckingRouterTest.java b/src/test/java/edu/kit/provideq/toolbox/api/tools/EquivalenceCheckingRouterTest.java new file mode 100644 index 00000000..8215c149 --- /dev/null +++ b/src/test/java/edu/kit/provideq/toolbox/api/tools/EquivalenceCheckingRouterTest.java @@ -0,0 +1,64 @@ +package edu.kit.provideq.toolbox.api.tools; + +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.when; + +import com.fasterxml.jackson.databind.JsonNode; +import com.fasterxml.jackson.databind.ObjectMapper; +import edu.kit.provideq.toolbox.tools.equivalencechecking.EquivalenceChecking; +import org.junit.jupiter.api.Test; +import org.springframework.test.web.reactive.server.WebTestClient; + +class EquivalenceCheckingRouterTest { + private final ObjectMapper objectMapper = new ObjectMapper(); + + @Test + void equivalenceCheckingReturnsJsonObject() throws Exception { + var equivalenceChecking = mock(EquivalenceChecking.class); + JsonNode toolOutput = objectMapper.readTree(""" + { + "strategy": "pyzx", + "status": "equivalent", + "globalPhaseIgnored": true + } + """); + when(equivalenceChecking.check(any(JsonNode.class))).thenReturn(toolOutput); + + var router = new EquivalenceCheckingRouter(equivalenceChecking).getEquivalenceCheckingRoutes(); + var client = WebTestClient.bindToRouterFunction(router).build(); + + client.post() + .uri(EquivalenceCheckingRouter.EQUIVALENCE_CHECKING_PATH) + .header("Accept", "application/json") + .header("Content-Type", "application/json") + .bodyValue(""" + { + "strategy": "pyzx", + "qasmA": "OPENQASM 2.0;", + "qasmB": "OPENQASM 2.0;" + } + """) + .exchange() + .expectStatus().isOk() + .expectHeader().contentType("application/json") + .expectBody() + .jsonPath("$.strategy").isEqualTo("pyzx") + .jsonPath("$.status").isEqualTo("equivalent") + .jsonPath("$.globalPhaseIgnored").isEqualTo(true); + } + + @Test + void equivalenceCheckingRejectsMissingBody() { + var router = new EquivalenceCheckingRouter( + mock(EquivalenceChecking.class)).getEquivalenceCheckingRoutes(); + var client = WebTestClient.bindToRouterFunction(router).build(); + + client.post() + .uri(EquivalenceCheckingRouter.EQUIVALENCE_CHECKING_PATH) + .header("Accept", "application/json") + .header("Content-Type", "application/json") + .exchange() + .expectStatus().isBadRequest(); + } +} diff --git a/src/test/java/edu/kit/provideq/toolbox/circuit/processing/results/ExecutionResultHelperTest.java b/src/test/java/edu/kit/provideq/toolbox/circuit/processing/results/ExecutionResultHelperTest.java new file mode 100644 index 00000000..db1e4c8a --- /dev/null +++ b/src/test/java/edu/kit/provideq/toolbox/circuit/processing/results/ExecutionResultHelperTest.java @@ -0,0 +1,161 @@ +package edu.kit.provideq.toolbox.circuit.processing.results; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.tuple; +import static org.assertj.core.api.Assertions.within; + +import edu.kit.provideq.toolbox.util.Pair; +import java.util.List; +import java.util.Optional; +import org.junit.jupiter.api.Test; + +class ExecutionResultHelperTest { + private static final double EPSILON = 1e-12; + + @Test + void createExecutionResult_parsesMultipleMeasurementsAndKeepsOriginalOptionals() { + ExecutionResult result = ExecutionResultHelper.createExecutionResult( + Optional.of("Counter({(1, 1, 0, 0): 159, (0, 0, 0, 0): 133, (0, 0, 1, 0): 101})"), + Optional.of("test-circuit") + ); + + assertThat(result.resultString()) + .contains("Counter({(1, 1, 0, 0): 159, (0, 0, 0, 0): 133, (0, 0, 1, 0): 101})"); + assertThat(result.circuit()).contains("test-circuit"); + + assertThat(result.sortedMeasurements()).isPresent(); + assertThat(result.sortedMeasurements().orElseThrow()) + .extracting(Pair::first, Pair::second) + .containsExactly( + tuple("1100", 159), + tuple("0000", 133), + tuple("0010", 101) + ); + } + + @Test + void createExecutionResult_calculatesProbabilitiesCorrectly() { + ExecutionResult result = ExecutionResultHelper.createExecutionResult( + Optional.of("Counter({(1, 1, 0, 0): 159, (0, 0, 0, 0): 133, (0, 0, 1, 0): 101})"), + Optional.empty() + ); + + List> probabilities = result.sortedProbabilities().orElseThrow(); + + assertThat(probabilities) + .extracting(Pair::first) + .containsExactly("1100", "0000", "0010"); + + int total = 159 + 133 + 101; + + assertThat(probabilities.get(0).second()).isCloseTo(159.0 / total, within(EPSILON)); + assertThat(probabilities.get(1).second()).isCloseTo(133.0 / total, within(EPSILON)); + assertThat(probabilities.get(2).second()).isCloseTo(101.0 / total, within(EPSILON)); + } + + @Test + void createExecutionResult_sortsDescendingIndependentOfInputOrder() { + ExecutionResult result = ExecutionResultHelper.createExecutionResult( + Optional.of("Counter({(0, 0): 5, (1, 1): 20, (1, 0): 10, (0, 1): 15})"), + Optional.empty() + ); + + assertThat(result.sortedMeasurements().orElseThrow()) + .extracting(Pair::first, Pair::second) + .containsExactly( + tuple("11", 20), + tuple("01", 15), + tuple("10", 10), + tuple("00", 5) + ); + + assertThat(result.sortedProbabilities().orElseThrow()) + .extracting(Pair::first) + .containsExactly("11", "01", "10", "00"); + } + + @Test + void createExecutionResult_handlesSingleMeasurement() { + ExecutionResult result = ExecutionResultHelper.createExecutionResult( + Optional.of("Counter({(1, 0, 1): 42})"), + Optional.empty() + ); + + assertThat(result.sortedMeasurements().orElseThrow()) + .extracting(Pair::first, Pair::second) + .containsExactly(tuple("101", 42)); + + assertThat(result.sortedProbabilities().orElseThrow()) + .extracting(Pair::first, Pair::second) + .containsExactly(tuple("101", 1.0)); + } + + @Test + void createExecutionResult_returnsEmptyListsForEmptyCounter_V1() { + ExecutionResult result = ExecutionResultHelper.createExecutionResult( + Optional.of("Counter()"), + Optional.empty() + ); + + assertThat(result.sortedMeasurements()).contains(List.of()); + assertThat(result.sortedProbabilities()).contains(List.of()); + } + + @Test + void createExecutionResult_returnsEmptyListsForEmptyCounter_V2() { + ExecutionResult result = ExecutionResultHelper.createExecutionResult( + Optional.of("Counter({})"), + Optional.empty() + ); + + assertThat(result.sortedMeasurements()).contains(List.of()); + assertThat(result.sortedProbabilities()).contains(List.of()); + } + + @Test + void createExecutionResult_keepsAllMeasurementsWhenCountsAreTied() { + ExecutionResult result = ExecutionResultHelper.createExecutionResult( + Optional.of("Counter({(1, 0): 7, (0, 1): 7, (1, 1): 3})"), + Optional.empty() + ); + + assertThat(result.sortedMeasurements().orElseThrow()) + .extracting(Pair::first, Pair::second) + .containsExactlyInAnyOrder( + tuple("10", 7), + tuple("01", 7), + tuple("11", 3) + ); + + assertThat(result.sortedMeasurements().orElseThrow()) + .extracting(Pair::second) + .containsExactly(7, 7, 3); + } + + @Test + void createExecutionResult_probabilitiesSumToOne() { + ExecutionResult result = ExecutionResultHelper.createExecutionResult( + Optional.of("Counter({(0, 0): 5, (0, 1): 15, (1, 0): 30})"), + Optional.empty() + ); + + double sum = result.sortedProbabilities().orElseThrow().stream() + .mapToDouble(Pair::second) + .sum(); + + assertThat(sum).isCloseTo(1.0, within(EPSILON)); + } + + @Test + void createExecutionResult_returnsEmptyMeasurementAndProbabilityOptionalsWhenResultStringIsEmpty() { + ExecutionResult result = ExecutionResultHelper.createExecutionResult( + Optional.empty(), + Optional.of("circuit") + ); + + assertThat(result.resultString()).isEmpty(); + assertThat(result.circuit()).contains("circuit"); + assertThat(result.sortedMeasurements()).isEmpty(); + assertThat(result.sortedProbabilities()).isEmpty(); + } +} diff --git a/src/test/java/edu/kit/provideq/toolbox/util/PairTest.java b/src/test/java/edu/kit/provideq/toolbox/util/PairTest.java new file mode 100644 index 00000000..8bff994e --- /dev/null +++ b/src/test/java/edu/kit/provideq/toolbox/util/PairTest.java @@ -0,0 +1,55 @@ +package edu.kit.provideq.toolbox.util; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertNotSame; + +import org.junit.jupiter.api.Test; + +class PairTest { + + @Test + void testPairConstructorAndAccessors() { + Pair pair = new Pair<>("test", 42); + + assertEquals("test", pair.first()); + assertEquals(42, pair.second()); + } + + @Test + void testCopyWithFirst() { + Pair original = new Pair<>("original", 10); + Pair updated = original.copyWithFirst("updated"); + + assertEquals("updated", updated.first()); + assertEquals(10, updated.second()); + assertEquals("original", original.first()); + assertNotSame(original, updated); + } + + @Test + void testCopyWithSecond() { + Pair original = new Pair<>("test", 10); + Pair updated = original.copyWithSecond(20); + + assertEquals("test", updated.first()); + assertEquals(20, updated.second()); + assertEquals(10, original.second()); + assertNotSame(original, updated); + } + + @Test + void testPairWithDifferentTypes() { + Pair pair = new Pair<>(3.14, "pi"); + + assertEquals(3.14, pair.first()); + assertEquals("pi", pair.second()); + + Pair updatedFirst = pair.copyWithFirst(2.71); + assertEquals(2.71, updatedFirst.first()); + assertEquals("pi", updatedFirst.second()); + + Pair updatedSecond = pair.copyWithSecond("e"); + assertEquals(3.14, updatedSecond.first()); + assertEquals("e", updatedSecond.second()); + } +}