dense-evolution 8.3.0__py3-none-win_amd64.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (165) hide show
  1. dashboard_core/__init__.py +115 -0
  2. dashboard_core/_gate_tables.py +30 -0
  3. dashboard_core/band_structure.py +71 -0
  4. dashboard_core/circuit_builder_component.py +232 -0
  5. dashboard_core/circuit_diagram.py +216 -0
  6. dashboard_core/crypto_protocols.py +77 -0
  7. dashboard_core/engine.py +326 -0
  8. dashboard_core/graphical_builder.py +114 -0
  9. dashboard_core/hamiltonians.py +593 -0
  10. dashboard_core/mass_decomposition_tool.py +47 -0
  11. dashboard_core/mitigation.py +343 -0
  12. dashboard_core/native_hf_diagnostics.py +62 -0
  13. dashboard_core/noise_tools.py +125 -0
  14. dashboard_core/qasm_library.py +233 -0
  15. dashboard_core/qmmm.py +16 -0
  16. dashboard_core/rag_tool.py +45 -0
  17. dashboard_core/state_visuals.py +288 -0
  18. dashboard_core/system_limits.py +60 -0
  19. dashboard_core/vector_healing.py +102 -0
  20. dashboard_core/visuals.py +158 -0
  21. dashboard_core/vqe.py +533 -0
  22. dashboard_core/wormhole.py +580 -0
  23. dense_evolution/__init__.py +114 -0
  24. dense_evolution/autodiff.py +10 -0
  25. dense_evolution/backends/__init__.py +5 -0
  26. dense_evolution/backends/chunk/__init__.py +37 -0
  27. dense_evolution/backends/chunk/_engine_imports.py +57 -0
  28. dense_evolution/backends/chunk/circuit_chunker.py +55 -0
  29. dense_evolution/backends/chunk/core.py +432 -0
  30. dense_evolution/backends/chunk/disk_overflow.py +232 -0
  31. dense_evolution/backends/chunk/geometry.py +95 -0
  32. dense_evolution/backends/chunk/guard.py +190 -0
  33. dense_evolution/backends/chunk/kernels.py +531 -0
  34. dense_evolution/backends/mps.py +1569 -0
  35. dense_evolution/backends/statevector.py +616 -0
  36. dense_evolution/chunk.py +25 -0
  37. dense_evolution/circuits/__init__.py +20 -0
  38. dense_evolution/circuits/compiler.py +488 -0
  39. dense_evolution/circuits/diagram.py +94 -0
  40. dense_evolution/circuits/gates.py +91 -0
  41. dense_evolution/circuits/parser.py +632 -0
  42. dense_evolution/circuits/qft.py +66 -0
  43. dense_evolution/circuits/random_circuit.py +85 -0
  44. dense_evolution/circuits/registry.py +74 -0
  45. dense_evolution/circuits/topology.py +79 -0
  46. dense_evolution/circuits/trotter.py +265 -0
  47. dense_evolution/circuits/uccsd.py +275 -0
  48. dense_evolution/cli.py +199 -0
  49. dense_evolution/compiler.py +9 -0
  50. dense_evolution/config.py +49 -0
  51. dense_evolution/drawing.py +10 -0
  52. dense_evolution/entropy.py +9 -0
  53. dense_evolution/fermions.py +9 -0
  54. dense_evolution/gates.py +9 -0
  55. dense_evolution/harrison_tb.py +16 -0
  56. dense_evolution/healing.py +18 -0
  57. dense_evolution/interop/__init__.py +18 -0
  58. dense_evolution/interop/qiskit_pennylane.py +406 -0
  59. dense_evolution/measurement.py +10 -0
  60. dense_evolution/mitigation/__init__.py +54 -0
  61. dense_evolution/mitigation/healing.py +215 -0
  62. dense_evolution/mitigation/kl_divergence.py +93 -0
  63. dense_evolution/mitigation/magic_entropy.py +163 -0
  64. dense_evolution/mitigation/magic_entropy_shadows.py +262 -0
  65. dense_evolution/mitigation/renyi.py +168 -0
  66. dense_evolution/mitigation/stabilizer_renyi_entropy.py +103 -0
  67. dense_evolution/mitigation/zne.py +990 -0
  68. dense_evolution/mps.py +9 -0
  69. dense_evolution/native_hf/__init__.py +26 -0
  70. dense_evolution/native_hf/_libcint/LICENSE-libcint +10 -0
  71. dense_evolution/native_hf/_libcint/libdecint.dll +0 -0
  72. dense_evolution/native_hf/assembly.py +304 -0
  73. dense_evolution/native_hf/basis.py +117 -0
  74. dense_evolution/native_hf/boys.py +35 -0
  75. dense_evolution/native_hf/bridge.py +112 -0
  76. dense_evolution/native_hf/cartesian.py +64 -0
  77. dense_evolution/native_hf/coulomb.py +196 -0
  78. dense_evolution/native_hf/differentiable.py +53 -0
  79. dense_evolution/native_hf/gaussians.py +79 -0
  80. dense_evolution/native_hf/kinetic.py +52 -0
  81. dense_evolution/native_hf/libcint_bridge.py +167 -0
  82. dense_evolution/native_hf/overlap.py +91 -0
  83. dense_evolution/native_hf/scf.py +404 -0
  84. dense_evolution/noise/__init__.py +79 -0
  85. dense_evolution/noise/coherent_attack.py +264 -0
  86. dense_evolution/noise/cosmic_ray.py +61 -0
  87. dense_evolution/noise/density_matrix_channels.py +78 -0
  88. dense_evolution/noise/differentiable.py +66 -0
  89. dense_evolution/noise/kraus/__init__.py +6 -0
  90. dense_evolution/noise/kraus/amplitude_damping.py +47 -0
  91. dense_evolution/noise/kraus/bitflip.py +22 -0
  92. dense_evolution/noise/kraus/combined.py +16 -0
  93. dense_evolution/noise/kraus/depolarizing.py +47 -0
  94. dense_evolution/noise/kraus/ideal.py +10 -0
  95. dense_evolution/noise/kraus/phaseflip.py +21 -0
  96. dense_evolution/noise/kraus_channels.py +285 -0
  97. dense_evolution/noise/oscillating.py +32 -0
  98. dense_evolution/noise/pink.py +80 -0
  99. dense_evolution/observables.py +11 -0
  100. dense_evolution/parser.py +9 -0
  101. dense_evolution/physics/__init__.py +27 -0
  102. dense_evolution/physics/entropy.py +161 -0
  103. dense_evolution/physics/fermions.py +322 -0
  104. dense_evolution/physics/observables.py +523 -0
  105. dense_evolution/physics/qec.py +1113 -0
  106. dense_evolution/physics/spectral.py +143 -0
  107. dense_evolution/physics/states.py +43 -0
  108. dense_evolution/protocols/__init__.py +27 -0
  109. dense_evolution/protocols/bb84.py +133 -0
  110. dense_evolution/protocols/di_qkd_ghz.py +199 -0
  111. dense_evolution/protocols/dicka_protocol2.py +124 -0
  112. dense_evolution/qec.py +20 -0
  113. dense_evolution/qft.py +9 -0
  114. dense_evolution/qmmm/__init__.py +13 -0
  115. dense_evolution/qmmm/ase_bridge.py +97 -0
  116. dense_evolution/qmmm/forces.py +388 -0
  117. dense_evolution/qmmm/propagation.py +80 -0
  118. dense_evolution/qmmm/region.py +137 -0
  119. dense_evolution/random_circuit.py +15 -0
  120. dense_evolution/registry.py +9 -0
  121. dense_evolution/simulator.py +10 -0
  122. dense_evolution/solvers/__init__.py +19 -0
  123. dense_evolution/solvers/autodiff.py +169 -0
  124. dense_evolution/solvers/harrison_tb.py +189 -0
  125. dense_evolution/solvers/vhd_tb.py +187 -0
  126. dense_evolution/states.py +9 -0
  127. dense_evolution/topology.py +9 -0
  128. dense_evolution/trotter.py +9 -0
  129. dense_evolution/utils/__init__.py +13 -0
  130. dense_evolution/utils/drawing.py +101 -0
  131. dense_evolution/utils/mass_decomposition.py +246 -0
  132. dense_evolution/utils/measurement.py +94 -0
  133. dense_evolution/vhd_tb.py +16 -0
  134. dense_evolution-8.3.0.dist-info/METADATA +366 -0
  135. dense_evolution-8.3.0.dist-info/RECORD +165 -0
  136. dense_evolution-8.3.0.dist-info/WHEEL +5 -0
  137. dense_evolution-8.3.0.dist-info/entry_points.txt +2 -0
  138. dense_evolution-8.3.0.dist-info/licenses/license.md +58 -0
  139. dense_evolution-8.3.0.dist-info/top_level.txt +5 -0
  140. ia_utils/__init__.py +0 -0
  141. ia_utils/adversarial_vector_attack.py +196 -0
  142. ia_utils/rag.py +288 -0
  143. ia_utils/vector_healing.py +399 -0
  144. local_site/__init__.py +0 -0
  145. local_site/app/__init__.py +0 -0
  146. local_site/app/server.py +1009 -0
  147. mcp_server/__init__.py +0 -0
  148. mcp_server/client.py +324 -0
  149. mcp_server/config.py +32 -0
  150. mcp_server/models.py +347 -0
  151. mcp_server/molecules.py +71 -0
  152. mcp_server/server.py +119 -0
  153. mcp_server/tools/__init__.py +0 -0
  154. mcp_server/tools/chemistry_tools.py +225 -0
  155. mcp_server/tools/circuit_tools.py +83 -0
  156. mcp_server/tools/crypto_tools.py +66 -0
  157. mcp_server/tools/mitigation_tools.py +81 -0
  158. mcp_server/tools/noise_tools.py +60 -0
  159. mcp_server/tools/retrieval_tools.py +44 -0
  160. mcp_server/tools/system_tools.py +149 -0
  161. mcp_server/tools/wormhole_tools.py +142 -0
  162. mcp_server/utils/__init__.py +0 -0
  163. mcp_server/utils/cache.py +55 -0
  164. mcp_server/utils/images.py +67 -0
  165. mcp_server/utils/truncation.py +38 -0
@@ -0,0 +1,616 @@
1
+ import warnings
2
+ import numpy as np
3
+ from typing import List, Tuple, Optional
4
+ from ..circuits.registry import HAS_JAX
5
+ from ..circuits.gates import GATES, PARAMETRIC_GATES, GATE_IDS, _TWO_QUBIT_PARAMETRIC_GATES
6
+ from ..circuits.compiler import QuantumTranspiler
7
+ from ..config import ensure_x64
8
+
9
+ if HAS_JAX:
10
+ import jax
11
+ import jax.numpy as jnp
12
+ # No jax.config.update here -- see dense_evolution/config.py. x64 is
13
+ # enabled lazily in __init__, and only when use_float32=False actually
14
+ # needs it, not as a side effect of merely importing this module.
15
+ from ..circuits.compiler import _compile_and_run_circuit_jit, _compile_and_run_circuit_jit_donated
16
+
17
+
18
+ # ─────────────────────────────────────────────────────────────────────────────
19
+ # Internal helpers
20
+ # ─────────────────────────────────────────────────────────────────────────────
21
+
22
+ def _qubit_stride_pairs(n: int, qubit: int):
23
+ """
24
+ Return (stride, outer_step, inner_step) for the MSB-first statevector
25
+ convention used throughout this simulator.
26
+
27
+ In MSB-first ordering qubit 0 is the *most* significant bit, so:
28
+ physical_bit_position = n - 1 - qubit
29
+ stride = 1 << physical_bit_position
30
+ """
31
+ phys = n - 1 - qubit
32
+ stride = 1 << phys
33
+ return stride
34
+
35
+
36
+ def _cx_numpy(sv: np.ndarray, n: int, ctrl: int, tgt: int) -> np.ndarray:
37
+ """
38
+ Vectorised CX on a NumPy statevector.
39
+ No Python loops — uses strided index arithmetic.
40
+ """
41
+ dim = len(sv)
42
+ c_stride = 1 << (n - 1 - ctrl)
43
+ t_stride = 1 << (n - 1 - tgt)
44
+ all_i = np.arange(dim, dtype=np.intp)
45
+ # Select indices where ctrl bit == 1 and tgt bit == 0
46
+ mask = ((all_i & c_stride) != 0) & ((all_i & t_stride) == 0)
47
+ idx_0 = all_i[mask]
48
+ idx_1 = idx_0 | t_stride
49
+ sv = sv.copy()
50
+ sv[idx_0], sv[idx_1] = sv[idx_1].copy(), sv[idx_0].copy()
51
+ return sv
52
+
53
+
54
+ def _cz_numpy(sv: np.ndarray, n: int, ctrl: int, tgt: int) -> np.ndarray:
55
+ """Vectorised CZ on a NumPy statevector."""
56
+ dim = len(sv)
57
+ c_stride = 1 << (n - 1 - ctrl)
58
+ t_stride = 1 << (n - 1 - tgt)
59
+ all_i = np.arange(dim, dtype=np.intp)
60
+ mask = ((all_i & c_stride) != 0) & ((all_i & t_stride) != 0)
61
+ sv = sv.copy()
62
+ sv[mask] *= -1
63
+ return sv
64
+
65
+
66
+ # ─────────────────────────────────────────────────────────────────────────────
67
+ # DenseSVSimulator
68
+ # ─────────────────────────────────────────────────────────────────────────────
69
+
70
+ class DenseSVSimulator:
71
+ """
72
+ Dense statevector quantum circuit simulator.
73
+
74
+ Qubit ordering: MSB-first (qubit 0 is the most significant bit).
75
+ Backends: NumPy (CPU), JAX XLA JIT (CPU/GPU/TPU) -- GPU dispatch is
76
+ automatic whenever a CUDA-enabled jaxlib is installed and a GPU is
77
+ present (jax.devices() reports it); no flag on this class selects it,
78
+ JAX's own default-device placement does.
79
+
80
+ Parameters
81
+ ----------
82
+ n_qubits : number of qubits
83
+ use_float32: use complex64 instead of complex128
84
+ """
85
+
86
+ def __init__(self, n_qubits: int,
87
+ use_float32: bool = False):
88
+ if n_qubits < 1 or n_qubits > 34:
89
+ raise ValueError(f"n_qubits must be in [1, 34], got {n_qubits}")
90
+ if HAS_JAX and not use_float32:
91
+ ensure_x64()
92
+ self.n = n_qubits
93
+ self.dim = 1 << n_qubits # 2 ** n_qubits
94
+ self.use_float32 = use_float32
95
+ self.dtype = np.complex64 if use_float32 else np.complex128
96
+ self.xp = jnp if HAS_JAX else np
97
+ self._reset_sv()
98
+
99
+ # ── state initialisation ─────────────────────────────────────────
100
+
101
+ def _reset_sv(self):
102
+ """Allocate |0...0⟩ on the active backend."""
103
+ if HAS_JAX:
104
+ self.sv = jnp.zeros(self.dim, dtype=self.dtype).at[0].set(1.0)
105
+ else:
106
+ self.sv = np.zeros(self.dim, dtype=self.dtype)
107
+ self.sv[0] = 1.0
108
+
109
+ def set_initial_state(self, state: Optional[np.ndarray] = None):
110
+ """
111
+ Reset the simulator.
112
+
113
+ Parameters
114
+ ----------
115
+ state : optional complex array of length 2**n.
116
+ If None, resets to |0...0⟩.
117
+ The array is normalised automatically.
118
+ """
119
+ if state is None:
120
+ self._reset_sv()
121
+ return
122
+ state = np.asarray(state, dtype=self.dtype)
123
+ if state.shape != (self.dim,):
124
+ raise ValueError(
125
+ f"State vector length {len(state)} != 2**{self.n} = {self.dim}")
126
+ norm = np.linalg.norm(state)
127
+ if norm < 1e-12:
128
+ raise ValueError("Cannot set a zero-norm state vector")
129
+ state = state / norm
130
+ if HAS_JAX:
131
+ self.sv = jnp.array(state)
132
+ else:
133
+ self.sv = state.copy()
134
+
135
+ # Alias used by the VQE engine
136
+ def set_state(self, state: np.ndarray):
137
+ self.set_initial_state(state)
138
+
139
+ # ── normalisation ─────────────────────────────────────────────────
140
+
141
+ def normalize(self):
142
+ norm = float(self.xp.linalg.norm(self.sv))
143
+ if norm > 1e-12:
144
+ if HAS_JAX:
145
+ self.sv = self.sv / norm
146
+ else:
147
+ self.sv /= norm
148
+
149
+ def _check_qubit_range(self, qubit: float, context: str) -> None:
150
+ """Validate a qubit index before it reaches the JIT-compiled fast
151
+ paths (run_circuit_jit / run_batch_jit).
152
+
153
+ Those paths encode qubit indices as bit-shift amounts inside
154
+ jax.lax.scan/switch and never call apply_gate_1q/apply_gate_2q
155
+ (which already validate) — an out-of-range index there doesn't
156
+ raise, it silently corrupts the entire statevector to zero
157
+ (verified: a single gate on an out-of-range qubit on an otherwise
158
+ normalized state left get_probabilities().sum() == 0.0, no error).
159
+ """
160
+ qi = int(qubit)
161
+ if not 0 <= qi < self.n:
162
+ raise ValueError(
163
+ f"Qubit index {qi} out of range [0, {self.n}) in {context}")
164
+
165
+ # ── 1-qubit gate ──────────────────────────────────────────────────
166
+
167
+ def apply_gate_1q(self, gate: np.ndarray, qubit: int):
168
+ """
169
+ Apply a 2×2 unitary to *qubit* via tensor contraction.
170
+
171
+ Uses reshape + moveaxis + matmul — fully vectorised,
172
+ no Python loops, compatible with both NumPy and JAX.
173
+ """
174
+ if not 0 <= qubit < self.n:
175
+ raise ValueError(f"Qubit index {qubit} out of range [0, {self.n})")
176
+ gate = self.xp.array(gate, dtype=self.dtype)
177
+ sv_nd = self.sv.reshape([2] * self.n)
178
+ sv_moved = self.xp.moveaxis(sv_nd, qubit, -1) # qubit axis → last
179
+ flat_shape = (self.dim >> 1, 2)
180
+ # matmul: (dim/2, 2) @ (2, 2).T → (dim/2, 2)
181
+ result = self.xp.dot(sv_moved.reshape(flat_shape),
182
+ gate.T)
183
+ self.sv = self.xp.moveaxis(
184
+ result.reshape([2] * self.n), -1, qubit).ravel()
185
+
186
+ # ── 2-qubit gate ──────────────────────────────────────────────────
187
+
188
+ def apply_gate_2q(self, gate: np.ndarray, q1: int, q2: int):
189
+ """
190
+ Apply a 4×4 unitary to qubits (q1, q2) via tensor contraction.
191
+ """
192
+ if q1 == q2:
193
+ raise ValueError("Control and target qubits must differ")
194
+ if not (0 <= q1 < self.n and 0 <= q2 < self.n):
195
+ raise ValueError(f"Qubit indices ({q1},{q2}) out of range [0, {self.n})")
196
+ gate = self.xp.array(gate, dtype=self.dtype)
197
+ sv_nd = self.sv.reshape([2] * self.n)
198
+ sv_moved = self.xp.moveaxis(sv_nd, (q1, q2), (-2, -1))
199
+ flat_shape = (self.dim >> 2, 4)
200
+ result = self.xp.dot(sv_moved.reshape(flat_shape),
201
+ gate.reshape(4, 4).T)
202
+ self.sv = self.xp.moveaxis(
203
+ result.reshape([2] * self.n), (-2, -1), (q1, q2)).ravel()
204
+
205
+ # ── specialised 2-qubit gates ─────────────────────────────────────
206
+
207
+ def apply_cx(self, ctrl: int, tgt: int):
208
+ """
209
+ CX (CNOT) gate.
210
+
211
+ JAX path: matrix contraction via apply_gate_2q.
212
+ NumPy path: fully vectorised index swap — no Python loops.
213
+ """
214
+ if ctrl == tgt:
215
+ raise ValueError("Control and target qubits must differ")
216
+ if not (0 <= ctrl < self.n and 0 <= tgt < self.n):
217
+ raise ValueError(f"Qubit indices ({ctrl},{tgt}) out of range [0, {self.n})")
218
+ if HAS_JAX:
219
+ cx_mat = jnp.array([
220
+ [1, 0, 0, 0],
221
+ [0, 1, 0, 0],
222
+ [0, 0, 0, 1],
223
+ [0, 0, 1, 0],
224
+ ], dtype=self.dtype)
225
+ self.apply_gate_2q(cx_mat, ctrl, tgt)
226
+ else:
227
+ self.sv = _cx_numpy(np.array(self.sv), self.n, ctrl, tgt)
228
+
229
+ def apply_cz(self, ctrl: int, tgt: int):
230
+ """
231
+ CZ gate.
232
+
233
+ JAX path: matrix contraction via apply_gate_2q.
234
+ NumPy path: fully vectorised sign flip — no Python loops.
235
+ """
236
+ if ctrl == tgt:
237
+ raise ValueError("Control and target qubits must differ")
238
+ if not (0 <= ctrl < self.n and 0 <= tgt < self.n):
239
+ raise ValueError(f"Qubit indices ({ctrl},{tgt}) out of range [0, {self.n})")
240
+ if HAS_JAX:
241
+ cz_mat = jnp.array([
242
+ [1, 0, 0, 0],
243
+ [0, 1, 0, 0],
244
+ [0, 0, 1, 0],
245
+ [0, 0, 0, -1],
246
+ ], dtype=self.dtype)
247
+ self.apply_gate_2q(cz_mat, ctrl, tgt)
248
+ else:
249
+ self.sv = _cz_numpy(np.array(self.sv), self.n, ctrl, tgt)
250
+
251
+ def apply_rx(self, qubit: int, theta: float):
252
+ """Apply a parameterized RX gate using the active backend (NumPy/JAX)."""
253
+ cos, sin = self.xp.cos(theta / 2), self.xp.sin(theta / 2)
254
+ mat = self.xp.array([[cos, -1j * sin], [-1j * sin, cos]], dtype=self.dtype)
255
+ self.apply_gate_1q(mat, qubit)
256
+
257
+ def apply_ry(self, qubit: int, theta: float):
258
+ """Apply a parameterized RY gate using the active backend (NumPy/JAX)."""
259
+ cos, sin = self.xp.cos(theta / 2), self.xp.sin(theta / 2)
260
+ mat = self.xp.array([[cos, -sin], [sin, cos]], dtype=self.dtype)
261
+ self.apply_gate_1q(mat, qubit)
262
+
263
+ def apply_rz(self, qubit: int, theta: float):
264
+ """Apply a parameterized RZ gate using the active backend (NumPy/JAX)."""
265
+ exp_neg = self.xp.exp(-1j * theta / 2)
266
+ exp_pos = self.xp.exp(1j * theta / 2)
267
+ mat = self.xp.array([[exp_neg, 0.0], [0.0, exp_pos]], dtype=self.dtype)
268
+ self.apply_gate_1q(mat, qubit)
269
+
270
+ # ── measurement ───────────────────────────────────────────────────
271
+
272
+ def measure(self, qubit_idx: int, jax_key: Optional["jax.Array"] = None) -> int:
273
+ """
274
+ Projective measurement on *qubit_idx*.
275
+
276
+ Returns 0 or 1 and collapses the statevector.
277
+ Uses MSB-first physical bit index: phys = n - 1 - qubit_idx.
278
+
279
+ BUG FIX (original): the original NumPy collapse wrote
280
+ sv_reshaped[:, 1 if result == 0 else 0, :] = 0.0
281
+ which zeroed the *wrong* basis state (0 when result=1, 1 when result=0)
282
+ and never normalised the JAX path.
283
+
284
+ jax_key : optional JAX PRNGKey. When given, the random outcome is
285
+ drawn via jax.random.choice(jax_key, ...) instead of the
286
+ global `np.random.choice` -- explicit, seedable, and
287
+ independent of NumPy's global RNG state, matching
288
+ registry.NoiseModel.apply_to_sv's own jax_key convention
289
+ (see that function's docstring for why explicit keys, not
290
+ hidden per-instance state, are this codebase's convention
291
+ for JAX-side reproducibility). Default (None) keeps the
292
+ original np.random.choice behavior unchanged, on both
293
+ backends -- this measurement's own state-collapse still
294
+ does concrete Python branching either way (the `result`
295
+ drives which basis-state slot gets zeroed), so passing a
296
+ key makes the *outcome* reproducible, not this method
297
+ jax.jit-traceable.
298
+ """
299
+ if not 0 <= qubit_idx < self.n:
300
+ raise ValueError(
301
+ f"Qubit {qubit_idx} out of range [0, {self.n})")
302
+
303
+ # BUG FIX: the JAX branch below reshapes to a [2]*n tensor and
304
+ # moveaxis'd -- the exact same indexing scheme apply_gate_1q
305
+ # uses (`moveaxis(sv_nd, qubit, -1)`, qubit axis == qubit index
306
+ # directly, no conversion). This method's JAX branch was instead
307
+ # using `phys = n-1-qubit_idx` for that moveaxis -- correct for
308
+ # the *NumPy* branch below (genuinely different flat/stride
309
+ # arithmetic on a raveled array), but wrong for the JAX branch's
310
+ # reshape-based indexing, silently reading/collapsing the WRONG
311
+ # qubit's marginal whenever qubit_idx != n-1-qubit_idx. Verified
312
+ # directly: X on qubit 0 of a 2-qubit register, then measure(0),
313
+ # returned 0 instead of 1 before this fix (see tests/unit/test_simulator.py's
314
+ # TestMeasurement class).
315
+ phys = self.n - 1 - qubit_idx
316
+ stride = 1 << phys
317
+
318
+ # ── compute marginal probabilities ──────────────────────────
319
+ if HAS_JAX:
320
+ probs = jnp.abs(self.sv) ** 2
321
+ sv_nd = probs.reshape([2] * self.n)
322
+ mv = jnp.moveaxis(sv_nd, qubit_idx, 0)
323
+ prob_0 = float(jnp.sum(mv[0]))
324
+ prob_1 = float(jnp.sum(mv[1]))
325
+ else:
326
+ sv_res = self.sv.reshape(-1, 2, stride)
327
+ prob_0 = float(np.sum(np.abs(sv_res[:, 0, :]) ** 2))
328
+ prob_1 = float(np.sum(np.abs(sv_res[:, 1, :]) ** 2))
329
+
330
+ total = prob_0 + prob_1
331
+ if total < 1e-12:
332
+ raise RuntimeError("Statevector norm is zero — cannot measure")
333
+ prob_0 /= total
334
+ prob_1 /= total
335
+
336
+ if jax_key is not None:
337
+ if not HAS_JAX:
338
+ raise ValueError("measure(jax_key=...) requires JAX to be installed.")
339
+ result = int(jax.random.choice(jax_key, jnp.array([0, 1]), p=jnp.array([prob_0, prob_1])))
340
+ else:
341
+ result = int(np.random.choice([0, 1], p=[prob_0, prob_1]))
342
+
343
+ # ── collapse ────────────────────────────────────────────────
344
+ # Zero out the amplitudes corresponding to the *opposite* outcome.
345
+ zero_slot = 1 - result # if result=0, zero slot 1; if result=1, zero slot 0
346
+
347
+ if HAS_JAX:
348
+ sv_nd = self.sv.reshape([2] * self.n)
349
+ mv = jnp.moveaxis(sv_nd, qubit_idx, 0)
350
+ mv = mv.at[zero_slot].set(0.0 + 0j)
351
+ self.sv = jnp.moveaxis(mv, 0, qubit_idx).ravel()
352
+ else:
353
+ sv_res = self.sv.reshape(-1, 2, stride)
354
+ sv_res[:, zero_slot, :] = 0.0
355
+ self.sv = sv_res.ravel()
356
+
357
+ self.normalize()
358
+ return result
359
+
360
+ # ── circuit execution ─────────────────────────────────────────────
361
+
362
+ def run_circuit(self, circuit: List[Tuple], transpile: bool = True):
363
+ target = QuantumTranspiler.transpile(circuit) if transpile else circuit
364
+
365
+ # Auto-dispatch to the compiled path whenever every gate in this
366
+ # circuit (post-transpile) is supported there -- measured 6x+
367
+ # faster on realistic circuits (eager per-gate dispatch pays a
368
+ # Python<->JAX round trip per gate; the compiled path is one XLA
369
+ # call for the whole circuit). The point: a caller who has never
370
+ # heard of run_circuit_jit still gets it automatically, for any
371
+ # circuit it can actually run -- only falls through to the eager
372
+ # loop below for the few gates GATE_IDS doesn't cover yet (see
373
+ # circuits/gates.py's own GATE_IDS comment for exactly which).
374
+ if HAS_JAX and all(
375
+ (cmd[0].lower() if isinstance(cmd[0], str) else str(cmd[0]).lower()) in GATE_IDS
376
+ for cmd in target
377
+ ):
378
+ self.run_circuit_jit(target)
379
+ return
380
+
381
+ for cmd in target:
382
+ name = cmd[0].lower()
383
+ args = cmd[1:]
384
+
385
+ if name in GATES:
386
+ mat = self.xp.array(GATES[name], dtype=self.dtype)
387
+ if mat.shape == (2, 2):
388
+ self.apply_gate_1q(mat, int(args[0]))
389
+ else:
390
+ self.apply_gate_2q(mat, int(args[0]), int(args[1]))
391
+
392
+ elif name in PARAMETRIC_GATES:
393
+ # Dispatch by gate NAME, not arg count -- a 1-qubit gate
394
+ # with 2 params (u2: qubit,phi,lam) and a 2-qubit gate
395
+ # with 1 param (cp/crz: q1,q2,theta) both have len(args)
396
+ # == 3, so arg-count alone is ambiguous (see
397
+ # _TWO_QUBIT_PARAMETRIC_GATES's own comment in gates.py).
398
+ if name in _TWO_QUBIT_PARAMETRIC_GATES:
399
+ mat = self.xp.array(PARAMETRIC_GATES[name](*args[2:]), dtype=self.dtype)
400
+ self.apply_gate_2q(mat, int(args[0]), int(args[1]))
401
+ else:
402
+ mat = self.xp.array(PARAMETRIC_GATES[name](*args[1:]), dtype=self.dtype)
403
+ self.apply_gate_1q(mat, int(args[0]))
404
+
405
+ else:
406
+ raise ValueError(
407
+ f"unknown gate '{cmd[0]}' -- not in GATES or PARAMETRIC_GATES. "
408
+ f"A typo in a gate name used to be silently dropped from the "
409
+ f"circuit instead of raising (issue #4)."
410
+ )
411
+
412
+
413
+
414
+ def run_circuit_jit(self, circuit: List):
415
+
416
+ if not HAS_JAX:
417
+ return self.run_circuit(circuit)
418
+
419
+ target = QuantumTranspiler.transpile(circuit)
420
+ compiled_ops = []
421
+
422
+ for cmd in target:
423
+ name = cmd[0].lower() if isinstance(cmd[0], str) else str(cmd[0]).lower()
424
+ if name not in GATE_IDS:
425
+ raise ValueError(
426
+ f"unknown gate '{cmd[0]}' -- not in GATE_IDS. A typo in a gate "
427
+ f"name used to be silently dropped from the circuit instead of "
428
+ f"raising (issue #4)."
429
+ )
430
+
431
+ g_id = float(GATE_IDS[name])
432
+ args = cmd[1:]
433
+
434
+ # ── gate argument parsing ──────────────────────────────
435
+ # 1-qubit parametric: (name, qubit, param)
436
+ if name in ('rx', 'ry', 'rz', 'p', 'u1', 'phase', 'gphase'):
437
+ q1 = float(args[0])
438
+ self._check_qubit_range(q1, f"gate '{name}'")
439
+ p = float(args[1]) if len(args) > 1 else 0.0
440
+ compiled_ops.append([g_id, q1, 0.0, p])
441
+
442
+ # 2-qubit parametric: (name, ctrl, tgt, param)
443
+ elif name in ('cp', 'crz', 'cphase'):
444
+ ctrl = float(args[0])
445
+ tgt = float(args[1]) if len(args) > 1 else 0.0
446
+ self._check_qubit_range(ctrl, f"gate '{name}' (control)")
447
+ self._check_qubit_range(tgt, f"gate '{name}' (target)")
448
+ p = float(args[2]) if len(args) > 2 else 0.0
449
+ compiled_ops.append([g_id, ctrl, tgt, p])
450
+
451
+ # 2-qubit non-parametric: (name, ctrl, tgt)
452
+ elif name in ('cx', 'cz', 'swap', 'cy'):
453
+ ctrl = float(args[0])
454
+ tgt = float(args[1]) if len(args) > 1 else 0.0
455
+ self._check_qubit_range(ctrl, f"gate '{name}' (control)")
456
+ self._check_qubit_range(tgt, f"gate '{name}' (target)")
457
+ compiled_ops.append([g_id, ctrl, tgt, 0.0])
458
+
459
+ # 1-qubit non-parametric: (name, qubit)
460
+ else:
461
+ q1 = float(args[0]) if args else 0.0
462
+ self._check_qubit_range(q1, f"gate '{name}'")
463
+ compiled_ops.append([g_id, q1, 0.0, 0.0])
464
+
465
+ if compiled_ops:
466
+ ops_jnp = jnp.array(compiled_ops, dtype=jnp.float64)
467
+ # Safe to donate self.sv here: it's rebound immediately below and
468
+ # no code path anywhere keeps a stale reference to the old buffer
469
+ # across this call (verified across chunked/repeated calls too,
470
+ # see _compile_and_run_circuit_jit_donated's docstring).
471
+ self.sv = _compile_and_run_circuit_jit_donated(self.sv, ops_jnp)
472
+
473
+ def run_circuit_jit_beast_mode(self, circuit: List):
474
+ """Deprecated alias for run_circuit_jit -- kept so code written
475
+ against any pre-8.1.46 release keeps working. Will be removed in
476
+ a future major version; switch to run_circuit_jit."""
477
+ warnings.warn(
478
+ "run_circuit_jit_beast_mode is deprecated, use run_circuit_jit instead "
479
+ "(same behavior, shorter name). This alias will be removed in a future release.",
480
+ DeprecationWarning, stacklevel=2,
481
+ )
482
+ return self.run_circuit_jit(circuit)
483
+
484
+ def run_circuit_with_chunking(self, circuit: List, chunk_size: int = 500):
485
+ """
486
+ Execute a circuit in chunks to avoid JIT recompilation on
487
+ large variable-length circuits.
488
+
489
+ Each chunk is a separate _compile_and_run_circuit_jit call
490
+ with a fixed-size ops array, allowing XLA to cache each size.
491
+ """
492
+ target = QuantumTranspiler.transpile(circuit)
493
+ for i in range(0, len(target), chunk_size):
494
+ self.run_circuit_jit(target[i: i + chunk_size])
495
+
496
+ def run_batch_jit(self,
497
+ base_circuit: List,
498
+ parameter_batch: np.ndarray) -> "jnp.ndarray":
499
+
500
+ if not HAS_JAX:
501
+ raise RuntimeError("run_batch_jit requires JAX")
502
+
503
+ target = QuantumTranspiler.transpile(base_circuit)
504
+ compiled_ops = []
505
+
506
+ for cmd in target:
507
+ name = cmd[0].lower() if isinstance(cmd[0], str) else str(cmd[0]).lower()
508
+ if name not in GATE_IDS:
509
+ raise ValueError(
510
+ f"unknown gate '{cmd[0]}' -- not in GATE_IDS. A typo in a gate "
511
+ f"name used to be silently dropped from the circuit instead of "
512
+ f"raising (issue #4)."
513
+ )
514
+ g_id = float(GATE_IDS[name])
515
+ args = cmd[1:]
516
+ if name in ('rx', 'ry', 'rz', 'p', 'u1', 'phase'):
517
+ q1 = float(args[0])
518
+ self._check_qubit_range(q1, f"gate '{name}'")
519
+ compiled_ops.append([g_id, q1, 0.0, -1.0]) # -1.0 = param slot
520
+ elif name in ('cp', 'crz', 'cphase'):
521
+ ctrl = float(args[0])
522
+ tgt = float(args[1]) if len(args) > 1 else 0.0
523
+ self._check_qubit_range(ctrl, f"gate '{name}' (control)")
524
+ self._check_qubit_range(tgt, f"gate '{name}' (target)")
525
+ compiled_ops.append([g_id, ctrl, tgt, -1.0])
526
+ elif name in ('cx', 'cz', 'swap', 'cy'):
527
+ ctrl = float(args[0])
528
+ tgt = float(args[1]) if len(args) > 1 else 0.0
529
+ self._check_qubit_range(ctrl, f"gate '{name}' (control)")
530
+ self._check_qubit_range(tgt, f"gate '{name}' (target)")
531
+ compiled_ops.append([g_id, ctrl, tgt, 0.0])
532
+ else:
533
+ q1 = float(args[0]) if args else 0.0
534
+ self._check_qubit_range(q1, f"gate '{name}'")
535
+ compiled_ops.append([g_id, q1, 0.0, 0.0])
536
+
537
+ n_param_slots = sum(1 for op in compiled_ops if op[3] == -1.0)
538
+ parameter_batch = np.asarray(parameter_batch)
539
+ if parameter_batch.ndim != 2 or parameter_batch.shape[1] != n_param_slots:
540
+ raise ValueError(
541
+ f"parameter_batch has {parameter_batch.shape[-1] if parameter_batch.ndim else 0} "
542
+ f"column(s) but base_circuit has {n_param_slots} parametric gate(s) (rx/ry/rz/p/u1/"
543
+ f"phase/cp/crz/cphase) -- one column per parametric gate, in gate-appearance order. "
544
+ f"A literal float passed for one of these gates is NOT exempt: it still consumes a "
545
+ f"positional column. A mismatch here used to be clipped silently by JAX's default "
546
+ f"out-of-bounds indexing instead of raising (issue #6)."
547
+ )
548
+
549
+ template = jnp.array(compiled_ops, dtype=jnp.float64)
550
+ # self.dtype, not a hardcoded jnp.complex128: _apply_gate_fast_step
551
+ # (compiler.py) already derives its working dtype from the input
552
+ # statevector specifically so use_float32=True instances run in
553
+ # complex64 end to end -- hardcoding complex128 here silently
554
+ # discarded that for every run_batch_jit/run_parametric_batch_jit
555
+ # call regardless of the instance's own configured dtype.
556
+ init_sv = jnp.zeros(self.dim, dtype=self.dtype).at[0].set(1.0)
557
+
558
+ def simulate_single_instance(single_params: "jnp.ndarray") -> "jnp.ndarray":
559
+ """Run one parameter vector through the circuit."""
560
+
561
+ def patch_and_apply(carry: "jnp.ndarray",
562
+ op: "jnp.ndarray"):
563
+ """
564
+ carry: jnp.int32 scalar — current parametric gate index.
565
+ op: [g_id, q1, q2, p_sentinel]
566
+ """
567
+ idx = carry
568
+ is_param = op[3] == -1.0
569
+ final_p = jnp.where(is_param, single_params[idx], op[3])
570
+ next_idx = jnp.where(is_param, idx + jnp.int32(1), idx)
571
+ patched = jnp.array([op[0], op[1], op[2], final_p],
572
+ dtype=jnp.float64)
573
+ return next_idx, patched
574
+
575
+ _, patched_ops = jax.lax.scan(
576
+ patch_and_apply,
577
+ jnp.int32(0), # BUG FIX: was (0,) tuple — must be a scalar
578
+ template,
579
+ )
580
+ return _compile_and_run_circuit_jit(init_sv, patched_ops)
581
+
582
+ return jax.jit(jax.vmap(simulate_single_instance, in_axes=(0,)))(
583
+ jnp.asarray(parameter_batch, dtype=jnp.float64)
584
+ )
585
+
586
+ def run_parametric_batch_jit(self, base_circuit: List, parameter_batch: np.ndarray) -> "jnp.ndarray":
587
+ """Deprecated alias for run_batch_jit -- kept so code written
588
+ against any pre-8.1.46 release keeps working. Will be removed in
589
+ a future major version; switch to run_batch_jit."""
590
+ warnings.warn(
591
+ "run_parametric_batch_jit is deprecated, use run_batch_jit instead "
592
+ "(same behavior, shorter name). This alias will be removed in a future release.",
593
+ DeprecationWarning, stacklevel=2,
594
+ )
595
+ return self.run_batch_jit(base_circuit, parameter_batch)
596
+
597
+ # ── observables ───────────────────────────────────────────────────
598
+
599
+ def get_probabilities(self) -> np.ndarray:
600
+ """Return measurement probability distribution as a NumPy float64 array."""
601
+ probs = np.array(self.xp.abs(self.sv) ** 2, dtype=np.float64)
602
+ # guard against floating-point leakage outside [0, 1]
603
+ probs = np.clip(probs, 0.0, 1.0)
604
+ total = probs.sum()
605
+ if total > 1e-12:
606
+ probs /= total
607
+ return probs
608
+
609
+ def get_statevector(self) -> np.ndarray:
610
+ """Return the current statevector as a NumPy complex array."""
611
+ return np.array(self.sv, dtype=self.dtype)
612
+
613
+ def memory_mb(self) -> float:
614
+ """Statevector memory footprint in megabytes."""
615
+ bytes_per_element = 8 if self.use_float32 else 16 # complex64=8, complex128=16
616
+ return self.dim * bytes_per_element / 1_000_000
@@ -0,0 +1,25 @@
1
+ """Backward-compatibility shim -- the real implementation moved to
2
+ dense_evolution.backends.chunk as part of the Phase 2 subpackage split
3
+ (see prog.txt). chunk.py was the one module left behind at the package
4
+ root when the rest of the split happened (everything else -- simulator,
5
+ compiler, gates, trotter, qec, ... -- was already moved with its own
6
+ shim); this closes that gap.
7
+
8
+ Unlike trotter.py/qec.py's shims (which re-export a short, stable public
9
+ list), dense_evolution.chunk is imported directly by module path in many
10
+ places -- tests/unit/test_chunk.py, tools/dashboard/core/system_limits.py,
11
+ research/local_site/app/server.py -- including private helpers like
12
+ _compile_multi_chunk_ops, not just the public Chunk class. Re-exporting a
13
+ curated name list would silently drop one of those on the next internal
14
+ refactor, so instead this shim replaces itself in sys.modules with the
15
+ real module object: `dense_evolution.chunk` and
16
+ `dense_evolution.backends.chunk` become the exact same module, byte for
17
+ byte, not two objects kept in sync by hand.
18
+
19
+ Import from dense_evolution.backends.chunk directly in new code.
20
+ """
21
+ import sys as _sys
22
+
23
+ from dense_evolution.backends import chunk as _real_chunk
24
+
25
+ _sys.modules[__name__] = _real_chunk
@@ -0,0 +1,20 @@
1
+ """Circuits subpackage: gate registry, parsing, compilation, topology."""
2
+ from .gates import GATES, PARAMETRIC_GATES, GATE_IDS
3
+ from .parser import QASMParser, QASMCircuit
4
+ from .compiler import QuantumTranspiler
5
+ from .registry import HAS_JAX, NoiseModel, NoiseSpec, QuantumHardwareRegistry
6
+ from .topology import entangling_layer, VALID_PATTERNS
7
+ from .qft import qft
8
+ from .random_circuit import random_circuit
9
+ from .trotter import pauli_rotation_ops, trotter_evolve_ops
10
+
11
+ __all__ = [
12
+ "GATES", "PARAMETRIC_GATES", "GATE_IDS",
13
+ "QASMParser", "QASMCircuit",
14
+ "QuantumTranspiler",
15
+ "HAS_JAX", "NoiseModel", "NoiseSpec", "QuantumHardwareRegistry",
16
+ "entangling_layer", "VALID_PATTERNS",
17
+ "qft",
18
+ "random_circuit",
19
+ "pauli_rotation_ops", "trotter_evolve_ops",
20
+ ]