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,1569 @@
1
+ """
2
+ MPSSimulator - Matrix Product State statevector simulator, JAX-backed.
3
+
4
+ Ported from the "TurboQuant TUREQ MPSSimulator v8.2 MatryoshkaFlash"
5
+ prototype (private research notebook, never published as part of the
6
+ dense-evolution package). Two real bugs were found and fixed by
7
+ independent verification against DenseSVSimulator before this module
8
+ existed in its current form:
9
+
10
+ 1. The original applied Lloyd-Max quantization to the SVD singular
11
+ values on every truncation ("PolarQuantizer"). Measured a real ~0.5%
12
+ Total Variation Distance error against DenseSVSimulator on an
13
+ 8-qubit entangling test circuit, with ZERO bond-dimension savings to
14
+ show for it. Dropped entirely -- this module keeps only the plain
15
+ adaptive SVD truncation (JSD-budget-driven bond dimension, the
16
+ author's own stopping criterion -- standard SVD truncation,
17
+ non-standard stopping metric).
18
+
19
+ 2. get_top_k_probable_states (originally "_extract_top_k_paths") picked
20
+ a single "best" bond index via argmax at each step instead of
21
+ correctly summing over the bond dimension. Measured 0/8 correct
22
+ states against the exact contraction on the same test circuit,
23
+ values off by ~30x. Fixed by propagating the true partial-contraction
24
+ vector through each bond (matches exactly, to machine precision, on
25
+ every state it finds) -- but note it's a genuine greedy beam search,
26
+ not an exact top-k finder: recall of the true top states grows with
27
+ beam width k but isn't guaranteed complete for any fixed k.
28
+
29
+ Originally ported in plain numpy (matching the prototype), then
30
+ converted to jax.numpy so the core tensor contractions (einsum, SVD) run
31
+ on the same backend as the rest of dense_evolution instead of a second,
32
+ inconsistent numerics stack. Re-verified against DenseSVSimulator after
33
+ the conversion -- see test_mps.py.
34
+
35
+ Uses whatever jax_enable_x64 precision is currently active in the
36
+ process (does not toggle it itself) -- same convention as
37
+ DenseSVSimulator/Chunk, which rely on the caller (dashboard_core.py's
38
+ run_simulation) to set precision, since jax_enable_x64 is a process-wide
39
+ flag and toggling it locally would leak to unrelated code running later
40
+ in the same process.
41
+
42
+ For circuits with LOW entanglement (product states, GHZ/Bell-like chains,
43
+ shallow local circuits), the bond dimension stays small regardless of
44
+ qubit count, so this scales to hundreds of qubits where DenseSVSimulator
45
+ (or Chunk) cannot -- see get_probabilities_sampled and
46
+ get_top_k_probable_states, neither of which ever materializes a
47
+ (2**n,)-shaped array. For HIGHLY entangled circuits the bond dimension
48
+ grows and this degrades back toward the same exponential cost
49
+ DenseSVSimulator has -- it is not a universal replacement, it is
50
+ complementary.
51
+ """
52
+
53
+ import dataclasses
54
+ import heapq
55
+ import warnings
56
+ from functools import partial
57
+ from typing import List, Optional, Tuple
58
+
59
+ import jax
60
+ import jax.numpy as jnp
61
+ import numpy as np
62
+
63
+ from ..config import ensure_x64
64
+ from ..physics.observables import _normalize_terms
65
+
66
+
67
+ def _real_dtype_for(complex_dtype) -> jnp.dtype:
68
+ """The real dtype that pairs with a given complex dtype -- float64
69
+ for complex128, float32 for complex64. Diagnostic/bookkeeping arrays
70
+ (chi, jsd, truncation error, entanglement entropy, the compiled op
71
+ table) are real-valued and must follow whichever precision the
72
+ quantum state itself is running at, instead of a dtype literal that
73
+ JAX silently downcasts (with a warning) whenever x64 isn't active."""
74
+ return jnp.float64 if complex_dtype == jnp.complex128 else jnp.float32
75
+
76
+
77
+ def _jsd_vectors_jax(p: jnp.ndarray, q: jnp.ndarray) -> jnp.ndarray:
78
+ """Same Jensen-Shannon Distance as _jsd_vectors, but returns a jnp
79
+ scalar instead of a Python float -- the float() cast in _jsd_vectors
80
+ forces concretization, which is fine called eagerly but raises
81
+ ConcretizationTypeError under jax.vmap/jax.jit tracing (needed by
82
+ _vectorized_chi_search below). Same math, no behavior difference for
83
+ eager callers, who go through _jsd_vectors instead."""
84
+ eps = 1e-12
85
+ p_norm = p / (jnp.sum(p) + eps)
86
+ q_norm = q / (jnp.sum(q) + eps)
87
+ m = 0.5 * (p_norm + q_norm)
88
+
89
+ def _kl(a, b):
90
+ mask = a > eps
91
+ # jnp.where still traces both branches, but log(0) only ever
92
+ # feeds into a term multiplied by 0 through the mask -- safe,
93
+ # and avoids a Python-side branch on a traced value.
94
+ safe_log = jnp.where(mask, jnp.log2(jnp.where(mask, a, 1.0) / (b + eps)), 0.0)
95
+ return jnp.sum(jnp.where(mask, a * safe_log, 0.0))
96
+
97
+ js = 0.5 * (_kl(p_norm, m) + _kl(q_norm, m))
98
+ # p.shape[0] is a static Python int (shape, not a value) even for a
99
+ # traced array under vmap, so this branch is safe at trace time.
100
+ dim_factor = float(np.log10(p.shape[0]) / 2.0) if p.shape[0] > 1 else 1.0
101
+ return jnp.sqrt(jnp.clip(js * dim_factor, 0.0, 1.0))
102
+
103
+
104
+ def _jsd_vectors(p: jnp.ndarray, q: jnp.ndarray) -> float:
105
+ """Adaptive Jensen-Shannon Distance used to size the truncated bond
106
+ dimension: scales with log10(dim) so the same JSD budget stays
107
+ meaningful whether the local Hilbert space is small or large."""
108
+ return float(_jsd_vectors_jax(p, q))
109
+
110
+
111
+ def _vectorized_chi_search_jax(S: jnp.ndarray, eps: float, jsd_budget: float, max_bond: int):
112
+ """JIT/vmap-compatible core of _vectorized_chi_search -- no
113
+ float()/int()/bool() concretizing casts, returns (chi_new, jsd_val) as
114
+ traced jnp scalars instead of Python int/float, so it's usable from
115
+ inside a jax.lax.scan/jax.jit region (needed by _build_mps_runner's
116
+ step function). Same algorithm _vectorized_chi_search wraps and casts
117
+ to Python types for eager callers -- see that docstring for the
118
+ verification this replicates the real while loop exactly."""
119
+ n = S.shape[0]
120
+ max_possible = min(n, max_bond)
121
+ mask_above_eps = S > eps
122
+ initial_chi = jnp.clip(jnp.sum(mask_above_eps), 1, max_possible)
123
+
124
+ norm_full = jnp.sum(S ** 2) + 1e-15
125
+ p_full = (S ** 2) / norm_full
126
+ # Real bug (found via a real-machine bisection down to 5e34a1e, not
127
+ # guessed): under jax.jit, _jsd_vectors_jax's own internal eps=1e-12
128
+ # mask does not reliably zero JSD contributions from p_full's
129
+ # already-negligible tail (~1e-18, from squaring SVD noise-floor
130
+ # singular values ~1e-9) when batched via jax.vmap across multiple
131
+ # candidates at once -- eager execution of the exact same formula on
132
+ # the exact same array gives the correct (zero) JSD there, jit does
133
+ # not. Zeroing p_full's own negligible tail here, before it ever
134
+ # reaches the vmapped JSD computation, verified to fix this under
135
+ # jit (confirmed both eager and jit now agree) without changing any
136
+ # value large enough to matter physically.
137
+ p_full = jnp.where(p_full > 1e-12, p_full, 0.0)
138
+
139
+ candidates = jnp.arange(1, max_possible + 1)
140
+ idx = jnp.arange(n)
141
+ trunc_mask = idx[None, :] < candidates[:, None]
142
+ p_trunc_all = jnp.where(trunc_mask, p_full[None, :], 0.0)
143
+
144
+ jsd_all = jax.vmap(lambda pt: _jsd_vectors_jax(p_full, pt))(p_trunc_all)
145
+
146
+ valid_candidate = candidates >= initial_chi
147
+ satisfies = (jsd_all <= jsd_budget) & valid_candidate
148
+ any_satisfies = jnp.any(satisfies)
149
+ first_idx = jnp.argmax(satisfies)
150
+ chi_new = jnp.where(any_satisfies, candidates[first_idx], max_possible)
151
+ jsd_val = jnp.where(any_satisfies, jsd_all[first_idx], jsd_all[max_possible - 1])
152
+ return chi_new, jsd_val
153
+
154
+
155
+ def _vectorized_chi_search(S: jnp.ndarray, eps: float, jsd_budget: float, max_bond: int):
156
+ """JIT/vmap-compatible replacement for the Python `while` loop in
157
+ _svd_truncate (host-syncs every iteration today). Computes the JSD for
158
+ every candidate truncation size in one vectorized pass instead of
159
+ incrementing chi_new one at a time, then picks the smallest candidate
160
+ >= the eps-cutoff-based starting point that satisfies jsd_budget --
161
+ exactly replicating the while loop's search direction and stopping
162
+ condition (verified: 0 mismatches in chi_new/jsd_val against the real
163
+ while loop across 171 real _svd_truncate calls from actual entangled
164
+ circuits at 3 different (n_qubits, max_bond, jsd_budget) configurations,
165
+ including the budget-violation fallback path).
166
+
167
+ Does NOT assume JSD is monotonic in the candidate size -- it restricts
168
+ the search to candidates >= the same starting point the while loop
169
+ starts from and takes the first (smallest) one satisfying the budget,
170
+ same as the loop would find by incrementing, whether or not JSD happens
171
+ to be monotonic there.
172
+
173
+ Returns (chi_new, jsd_val) as Python int/float -- eager convenience
174
+ wrapper around _vectorized_chi_search_jax, for callers outside a jit
175
+ region (e.g. these same tests).
176
+ """
177
+ chi_new, jsd_val = _vectorized_chi_search_jax(S, eps, jsd_budget, max_bond)
178
+ return int(chi_new), float(jsd_val)
179
+
180
+
181
+ def _expand_nonlocal_2q_positions(q1: int, q2: int):
182
+ """Pure-Python prediction of the exact adjacent-pair call sequence
183
+ MPSSimulator._apply_nonlocal_2q's runtime SWAP-chain dispatch produces
184
+ for a given (q1, q2) -- moves that decision from runtime (inside the
185
+ eager gate-application loop) to pre-compile time, so it can be used to
186
+ build a flat op list before entering a JIT region, mirroring how
187
+ chunk.py's _compile_multi_chunk_ops does all gate-name branching
188
+ outside the traced kernel.
189
+
190
+ Verified exhaustively against the real runtime dispatch: 0 mismatches
191
+ across all (q1, q2) pairs for n_qubits in {6, 10, 15} (330 pairs).
192
+
193
+ Returns (seq, gate_index, needs_transpose):
194
+ seq -- list of (a, a+1) adjacent physical-position pairs,
195
+ in the exact order _apply_nonlocal_2q would call
196
+ apply_gate_2q on them.
197
+ gate_index -- which position in `seq` is the REAL gate (every
198
+ other position is a SWAP).
199
+ needs_transpose -- True iff the original q1 > q2, matching
200
+ _apply_nonlocal_2q's own q1>q2 normalization
201
+ (the gate tensor must be transposed via axes
202
+ (1, 0, 3, 2) before use at `seq[gate_index]`,
203
+ same as _apply_nonlocal_2q already does).
204
+ """
205
+ needs_transpose = q1 > q2
206
+ target, far = (q2, q1) if needs_transpose else (q1, q2)
207
+ seq = []
208
+ for q in range(far - 1, target, -1):
209
+ seq.append((q, q + 1))
210
+ gate_index = len(seq)
211
+ seq.append((target, target + 1))
212
+ for q in range(target + 1, far):
213
+ seq.append((q, q + 1))
214
+ return seq, gate_index, needs_transpose
215
+
216
+
217
+ # ── gate-ID -> matrix construction, jax.lax.switch-based (JIT-fusion) ────
218
+ # Same gate-ID table as dense_evolution/gates.py::GATE_IDS and
219
+ # compiler.py::_apply_gate_fast_step's g_1q switch -- same formulas,
220
+ # verified to match exactly (0 error, both eagerly and under jax.vmap/jit)
221
+ # before being committed here. Unlike compiler.py's statevector-bitmask
222
+ # do_2q, MPS's own gate application is an explicit tensor contraction
223
+ # (jnp.einsum in apply_gate_2q), so the 2-qubit switch below returns the
224
+ # full (2,2,2,2) unitary tensor directly -- same matrices mps.py's own
225
+ # apply_cx/apply_cz/apply_swap already build by hand, generalized to a
226
+ # traced g_id instead of one hardcoded gate per method. id=23 (SWAP) is
227
+ # reserved-but-unused in gates.py's own table (comment there: "never
228
+ # dispatched here, QuantumTranspiler always decomposes swap into 3xCX
229
+ # first") -- MPS's own kernel claims it for real, since decomposing SWAP
230
+ # into 3 CX would mean 3 SVD truncations per swap instead of 1, defeating
231
+ # the point of keeping SWAP as its own gate for the non-adjacent-gate
232
+ # chain (see _expand_nonlocal_2q_positions above).
233
+
234
+ def _mps_1q_matrix(g_id: jnp.ndarray, param: jnp.ndarray, dtype) -> jnp.ndarray:
235
+ """Traced gate-ID -> (2,2) unitary, ids 0-13 (I,H,X,Y,Z,S,Sdg,T,Tdg,
236
+ Rx,Ry,Rz,P,SX) -- same formulas as compiler.py's g_1q switch."""
237
+ inv2 = jnp.asarray(1.0 / jnp.sqrt(2.0), dtype=dtype)
238
+ half_p = param * 0.5
239
+ cos_p = jnp.cos(half_p).astype(dtype)
240
+ sin_p = jnp.sin(half_p).astype(dtype)
241
+ exp_pos = jnp.exp(1j * param).astype(dtype)
242
+ exp_ph4 = jnp.exp(1j * jnp.pi / 4.0).astype(dtype)
243
+ exp_mh4 = jnp.exp(-1j * jnp.pi / 4.0).astype(dtype)
244
+ safe_gid = jnp.clip(g_id, 0, 13)
245
+ return jax.lax.switch(
246
+ safe_gid,
247
+ [
248
+ lambda _: jnp.eye(2, dtype=dtype), # 0 I
249
+ lambda _: jnp.array([[inv2, inv2], [inv2, -inv2]], dtype=dtype), # 1 H
250
+ lambda _: jnp.array([[0.0 + 0j, 1.0 + 0j], [1.0 + 0j, 0.0 + 0j]], dtype=dtype), # 2 X
251
+ lambda _: jnp.array([[0.0 + 0j, -1j], [1j, 0.0 + 0j]], dtype=dtype), # 3 Y
252
+ lambda _: jnp.array([[1.0 + 0j, 0.0 + 0j], [0.0 + 0j, -1.0 + 0j]], dtype=dtype), # 4 Z
253
+ lambda _: jnp.array([[1.0 + 0j, 0.0 + 0j], [0.0 + 0j, 1j]], dtype=dtype), # 5 S
254
+ lambda _: jnp.array([[1.0 + 0j, 0.0 + 0j], [0.0 + 0j, -1j]], dtype=dtype), # 6 Sdg
255
+ lambda _: jnp.array([[1.0 + 0j, 0.0 + 0j], [0.0 + 0j, exp_ph4]], dtype=dtype), # 7 T
256
+ lambda _: jnp.array([[1.0 + 0j, 0.0 + 0j], [0.0 + 0j, exp_mh4]], dtype=dtype), # 8 Tdg
257
+ lambda _: jnp.array([[cos_p, -1j * sin_p], [-1j * sin_p, cos_p]], dtype=dtype), # 9 Rx
258
+ lambda _: jnp.array([[cos_p, -sin_p], [sin_p, cos_p]], dtype=dtype), # 10 Ry
259
+ lambda _: jnp.array([[jnp.exp(-1j * half_p), 0.0 + 0j],
260
+ [0.0 + 0j, jnp.exp(1j * half_p)]], dtype=dtype), # 11 Rz
261
+ lambda _: jnp.array([[1.0 + 0j, 0.0 + 0j], [0.0 + 0j, exp_pos]], dtype=dtype), # 12 P
262
+ lambda _: jnp.array([[0.5 + 0.5j, 0.5 - 0.5j], [0.5 - 0.5j, 0.5 + 0.5j]], dtype=dtype), # 13 SX
263
+ ],
264
+ operand=None,
265
+ )
266
+
267
+
268
+ def _mps_2q_matrix(g_id: jnp.ndarray, param: jnp.ndarray, dtype) -> jnp.ndarray:
269
+ """Traced gate-ID -> (2,2,2,2) unitary tensor, ids 20-25
270
+ (CX,CZ,CP,SWAP,CY,CRZ) -- same target gates as compiler.py's do_2q,
271
+ built as explicit tensors (not the statevector-bitmask approach) to
272
+ match MPS's own einsum-based 2-qubit gate application."""
273
+ exp_pos = jnp.exp(1j * param).astype(dtype)
274
+ exp_neg_half = jnp.exp(-1j * param * 0.5).astype(dtype)
275
+ exp_pos_half = jnp.exp(1j * param * 0.5).astype(dtype)
276
+ one = jnp.asarray(1.0 + 0j, dtype=dtype)
277
+ cp_diag = jnp.stack([one, one, one, exp_pos])
278
+ crz_diag = jnp.stack([one, one, exp_neg_half, exp_pos_half])
279
+ safe_idx = jnp.clip(g_id - 20, 0, 5)
280
+ mat = jax.lax.switch(
281
+ safe_idx,
282
+ [
283
+ lambda _: jnp.array([[1, 0, 0, 0], [0, 1, 0, 0], [0, 0, 0, 1], [0, 0, 1, 0]], dtype=dtype), # 20 CX
284
+ lambda _: jnp.diag(jnp.array([1, 1, 1, -1], dtype=dtype)), # 21 CZ
285
+ lambda _: jnp.diag(cp_diag), # 22 CP
286
+ lambda _: jnp.array([[1, 0, 0, 0], [0, 0, 1, 0], [0, 1, 0, 0], [0, 0, 0, 1]], dtype=dtype), # 23 SWAP
287
+ lambda _: jnp.array([[1, 0, 0, 0], [0, 1, 0, 0], [0, 0, 0, -1j], [0, 0, 1j, 0]], dtype=dtype), # 24 CY
288
+ lambda _: jnp.diag(crz_diag), # 25 CRZ
289
+ ],
290
+ operand=None,
291
+ )
292
+ return mat.reshape(2, 2, 2, 2)
293
+
294
+
295
+ def _pad_gamma(g: jnp.ndarray, max_bond: int) -> jnp.ndarray:
296
+ """Zero-pads a (chiL, 2, chiR) gamma tensor up to (max_bond, 2,
297
+ max_bond) -- the fixed shape run_circuit_jit's scan carry needs.
298
+ Mathematically transparent: every contraction used here (einsum) sums
299
+ over the padded axis, and zero entries there contribute exactly zero,
300
+ verified directly (max singular-value error at machine epsilon,
301
+ final-state fidelity 1.0 to 1e-13, against the real eager simulator
302
+ across multiple (n_qubits, max_bond, jsd_budget) configurations
303
+ including budget-violation cases)."""
304
+ chi_l, d, chi_r = g.shape
305
+ out = jnp.zeros((max_bond, d, max_bond), dtype=g.dtype)
306
+ return out.at[:chi_l, :, :chi_r].set(g)
307
+
308
+
309
+ def _pad_lambda(lam: jnp.ndarray, max_bond: int) -> jnp.ndarray:
310
+ """Zero-pads a (chi,) lambda vector up to (max_bond,)."""
311
+ n = lam.shape[0]
312
+ out = jnp.zeros((max_bond,), dtype=lam.dtype)
313
+ return out.at[:n].set(lam)
314
+
315
+
316
+ @partial(jax.jit, static_argnums=(1, 2))
317
+ def _pad_all_gammas(gammas, max_bond: int, dtype) -> jnp.ndarray:
318
+ """Batches the whole per-gamma padding+stacking loop into one
319
+ compiled call instead of n_qubits separate eager _pad_gamma
320
+ dispatches -- profiling run_circuit_jit directly found this eager
321
+ loop (50-51 separate JAX dispatches for N=50) cost 52-55% of total
322
+ run_circuit_jit time, present equally on the fuse_gates=True and
323
+ default paths, diluting the real speedup ratio between them.
324
+ Persistent-compiled-cache wrapper, keyed by the tuple of gamma
325
+ shapes + max_bond + dtype -- valid across the whole process
326
+ lifetime, same principle as _jit_mps_1q_matrix/_jit_mps_2q_matrix."""
327
+ return jnp.stack([_pad_gamma(g, max_bond).astype(dtype) for g in gammas])
328
+
329
+
330
+ @partial(jax.jit, static_argnums=(1, 2))
331
+ def _pad_all_lambdas(lambdas, max_bond: int, dtype) -> jnp.ndarray:
332
+ """Same persistent-compiled-cache fix as _pad_all_gammas, for
333
+ lambdas."""
334
+ return jnp.stack([_pad_lambda(l, max_bond).astype(dtype) for l in lambdas])
335
+
336
+
337
+ def _compile_mps_ops(ops, n_qubits: int) -> List[List[float]]:
338
+ """Pre-compile-time (pure Python, outside any jit region) translation
339
+ of a circuit -- list of (name, *args) tuples/lists, same convention as
340
+ DenseSVSimulator.run_circuit_jit_beast_mode -- into a flat list of
341
+ [g_id, q1, q2, param, transpose_flag] rows for run_circuit_jit's
342
+ jax.lax.scan kernel. Two things happen here, deliberately outside the
343
+ jit region (same principle as chunk.py's _compile_multi_chunk_ops --
344
+ all Python-level branching on gate identity happens before tracing
345
+ starts, never inside the traced step function):
346
+
347
+ 1. GATE_IDS (dense_evolution/gates.py) name -> numeric id lookup.
348
+ 2. Non-adjacent 2-qubit gates, and adjacent ones with q1 > q2,
349
+ expanded/normalized via _expand_nonlocal_2q_positions (verified
350
+ above) into their SWAP-chain-equivalent sequence of
351
+ adjacent-ascending rows -- moves what _apply_nonlocal_2q's
352
+ runtime dispatch does today to pre-compile time.
353
+
354
+ id 23 (SWAP) is used for real here (not decomposed into 3xCX like
355
+ QuantumTranspiler.transpile would) -- see _mps_2q_matrix's docstring
356
+ for why: decomposing would mean 3 SVD truncations per swap instead
357
+ of 1, defeating the point of the chain in the first place.
358
+ """
359
+ from ..circuits.compiler import QuantumTranspiler
360
+ from ..circuits.gates import GATE_IDS
361
+
362
+ # CCX has no entry in GATE_IDS (it's not a 1- or 2-qubit gate) -- reuse
363
+ # the same Barenco H/CX/T/Tdg decomposition compiler.py/chunk.py's own
364
+ # QuantumTranspiler already uses, rather than inventing a second one.
365
+ # 'swap' is deliberately NOT pre-decomposed here (see id 23 note below).
366
+ expanded_ops = []
367
+ for cmd in ops:
368
+ if str(cmd[0]).lower() == 'ccx':
369
+ expanded_ops.extend(QuantumTranspiler.decompose_toffoli(*cmd[1:4]))
370
+ else:
371
+ expanded_ops.append(cmd)
372
+
373
+ rows: List[List[float]] = []
374
+ for cmd in expanded_ops:
375
+ name = str(cmd[0]).lower()
376
+ if name not in GATE_IDS:
377
+ raise ValueError(
378
+ f"unknown gate '{cmd[0]}' -- not in GATE_IDS. run_circuit_jit "
379
+ f"does not silently drop unrecognized gates (same policy as "
380
+ f"run_circuit_jit_beast_mode, issue #4)."
381
+ )
382
+ g_id = float(GATE_IDS[name])
383
+
384
+ if g_id < 20:
385
+ q1 = int(cmd[1])
386
+ param = float(cmd[2]) if len(cmd) > 2 else 0.0
387
+ rows.append([g_id, float(q1), float(q1), param, 0.0])
388
+ continue
389
+
390
+ q1 = int(cmd[1])
391
+ q2 = int(cmd[2])
392
+ param = float(cmd[3]) if len(cmd) > 3 else 0.0
393
+
394
+ if abs(q1 - q2) == 1:
395
+ lo, hi = (q1, q2) if q1 < q2 else (q2, q1)
396
+ needs_transpose = q1 > q2
397
+ rows.append([g_id, float(lo), float(hi), param, 1.0 if needs_transpose else 0.0])
398
+ else:
399
+ seq, gate_index, needs_transpose = _expand_nonlocal_2q_positions(q1, q2)
400
+ for i, (a, b) in enumerate(seq):
401
+ if i == gate_index:
402
+ rows.append([g_id, float(a), float(b), param, 1.0 if needs_transpose else 0.0])
403
+ else:
404
+ rows.append([23.0, float(a), float(b), 0.0, 0.0]) # SWAP
405
+
406
+ return rows
407
+
408
+
409
+ def _embed_1q_matrix(mat1: np.ndarray, qubit: int, pair: Tuple[int, int]) -> np.ndarray:
410
+ """Embeds a 2x2 single-qubit matrix into the 4x4 space of a 2-qubit
411
+ pair (a, b), acting as identity on the other qubit -- kron(mat1, I2)
412
+ if qubit is the pair's first (more significant) qubit, kron(I2, mat1)
413
+ otherwise, matching _mps_2q_matrix's own convention (its (4,4)
414
+ matrices are built with qubit0 as the more significant index before
415
+ reshaping to (2,2,2,2))."""
416
+ eye2 = np.eye(2, dtype=mat1.dtype)
417
+ a, _ = pair
418
+ return np.kron(mat1, eye2) if qubit == a else np.kron(eye2, mat1)
419
+
420
+
421
+ @partial(jax.jit, static_argnums=(2,))
422
+ def _jit_mps_1q_matrix(g_id, param, dtype):
423
+ """Persistent-compiled-cache wrapper around _mps_1q_matrix, keyed by
424
+ JAX on (g_id, param, dtype) across the whole process lifetime, not
425
+ just within one _fuse_compiled_rows call -- without this, each
426
+ separate run_circuit_jit(fuse_gates=True) call pays a real XLA
427
+ compile cost again for the same gate/parameter combination, since
428
+ _mps_1q_matrix's own jax.lax.switch is never wrapped in @jax.jit."""
429
+ return _mps_1q_matrix(g_id, param, dtype)
430
+
431
+
432
+ @partial(jax.jit, static_argnums=(2,))
433
+ def _jit_mps_2q_matrix(g_id, param, dtype):
434
+ """Same persistent-compiled-cache fix as _jit_mps_1q_matrix, for
435
+ _mps_2q_matrix."""
436
+ return _mps_2q_matrix(g_id, param, dtype)
437
+
438
+
439
+ def _fuse_compiled_rows(rows: List[List[float]], dtype) -> List[tuple]:
440
+ """Host-side (pure Python/NumPy, outside any jit region), exact --
441
+ fuses a chain of consecutive _compile_mps_ops rows into one matrix
442
+ via matrix multiplication, growing the active qubit pair as needed:
443
+ a 1-qubit gate on a qubit already inside the current chain's pair
444
+ gets embedded into that same 4x4 space (kron with identity on the
445
+ other qubit) rather than breaking the chain -- this is what actually
446
+ fuses a CX-RZ-CX triple (RZ acts on only one of CX's two qubits, so a
447
+ naive "exact qubit set match" rule never fuses anything). A chain
448
+ starting on a single qubit that is later joined by a 2-qubit gate
449
+ touching it also grows into the 2-qubit space the same way. Any gate
450
+ touching a qubit OUTSIDE the current chain's pair ends the chain.
451
+
452
+ Runs on _compile_mps_ops's OWN output rows, so CCX decomposition and
453
+ SWAP-chain expansion for non-adjacent gates have already happened --
454
+ fusion composes correctly with both by construction (a SWAP is just
455
+ another 2-qubit gate, fusable like any other), verified directly
456
+ against the eager reference on both a non-adjacent-gate circuit and
457
+ a CCX circuit, not merely assumed.
458
+
459
+ Returns a list of ('1q', q, matrix_2x2) / ('2q', q1, q2, matrix_4x4)
460
+ entries. transpose_flag is resolved into the matrix itself here, so
461
+ the fused representation never needs it downstream.
462
+
463
+ Memoizes row_matrix by (g_id, param, transpose_flag) -- most rows in
464
+ a real circuit share the same small set of distinct gate/parameter
465
+ combinations (e.g. one N=50 TFIM circuit has 985 rows but only 3
466
+ distinct ones), and each cache miss dispatches a real JAX call
467
+ (_mps_1q_matrix/_mps_2q_matrix) that is not free to skip. Without
468
+ this, a real (not tiny-example) circuit pays a full eager JAX
469
+ dispatch per row on every single run_circuit_jit(fuse_gates=True)
470
+ call, not just the first -- measured directly: an unmemoized version
471
+ of this function made "warm" (post-compile) calls nearly as slow as
472
+ "cold" ones (10.7s vs 12.9s on a 74-gate circuit), because the
473
+ matrix-reconstruction cost, not recompilation, dominated."""
474
+ matrix_cache: dict = {}
475
+
476
+ def row_matrix(row):
477
+ g_id, q1, q2, param, transpose_flag = row
478
+ key = (g_id, param, transpose_flag)
479
+ cached = matrix_cache.get(key)
480
+ if cached is not None:
481
+ return cached, (int(q1), int(q2)) if g_id >= 20 else (int(q1),)
482
+ g_id_arr = jnp.asarray(g_id).astype(jnp.int32)
483
+ if g_id >= 20:
484
+ mat4 = np.asarray(_jit_mps_2q_matrix(g_id_arr, jnp.asarray(param), dtype))
485
+ if transpose_flag > 0.5:
486
+ mat4 = np.transpose(mat4, (1, 0, 3, 2))
487
+ mat = mat4.reshape(4, 4)
488
+ else:
489
+ mat = np.asarray(_jit_mps_1q_matrix(g_id_arr, jnp.asarray(param), dtype))
490
+ matrix_cache[key] = mat
491
+ return mat, (int(q1), int(q2)) if g_id >= 20 else (int(q1),)
492
+
493
+ fused = []
494
+ i = 0
495
+ n = len(rows)
496
+ while i < n:
497
+ mat, qubits = row_matrix(rows[i])
498
+ pair = qubits if len(qubits) == 2 else None
499
+ active_qubit = qubits[0]
500
+ j = i + 1
501
+ while j < n:
502
+ nmat, nqubits = row_matrix(rows[j])
503
+ n_is_2q = len(nqubits) == 2
504
+ if pair is None:
505
+ if not n_is_2q:
506
+ if nqubits[0] != active_qubit:
507
+ break
508
+ mat = nmat @ mat
509
+ j += 1
510
+ continue
511
+ if active_qubit not in nqubits:
512
+ break
513
+ pair = nqubits
514
+ mat = _embed_1q_matrix(mat, active_qubit, pair)
515
+ mat = nmat @ mat
516
+ j += 1
517
+ continue
518
+ if n_is_2q:
519
+ if nqubits != pair:
520
+ break
521
+ else:
522
+ if nqubits[0] not in pair:
523
+ break
524
+ nmat = _embed_1q_matrix(nmat, nqubits[0], pair)
525
+ mat = nmat @ mat
526
+ j += 1
527
+ if pair is not None:
528
+ fused.append(('2q', pair[0], pair[1], mat.reshape(2, 2, 2, 2)))
529
+ else:
530
+ fused.append(('1q', active_qubit, mat))
531
+ i = j
532
+ return fused
533
+
534
+
535
+ def _fused_entries_to_arrays(fused: List[tuple], dtype):
536
+ """Turns _fuse_compiled_rows's output into the parallel arrays
537
+ jax.lax.scan needs -- every step carries BOTH a 1q and a 2q matrix
538
+ slot (one is an unused identity placeholder) so shapes stay uniform
539
+ across the whole scanned sequence, same principle _mps_1q_matrix/
540
+ _mps_2q_matrix's own switch already relies on (all branches traced,
541
+ only one path's real work executed)."""
542
+ is_2q, q1s, q2s, mats2q, mats1q = [], [], [], [], []
543
+ eye2 = np.eye(2, dtype=dtype)
544
+ eye4 = np.eye(4, dtype=dtype).reshape(2, 2, 2, 2)
545
+ for entry in fused:
546
+ if entry[0] == '2q':
547
+ _, a, b, mat = entry
548
+ is_2q.append(True)
549
+ q1s.append(a)
550
+ q2s.append(b)
551
+ mats2q.append(mat)
552
+ mats1q.append(eye2)
553
+ else:
554
+ _, a, mat = entry
555
+ is_2q.append(False)
556
+ q1s.append(a)
557
+ q2s.append(0)
558
+ mats2q.append(eye4)
559
+ mats1q.append(mat)
560
+ return (jnp.asarray(is_2q), jnp.asarray(q1s, dtype=jnp.int32), jnp.asarray(q2s, dtype=jnp.int32),
561
+ jnp.asarray(np.stack(mats2q)), jnp.asarray(np.stack(mats1q)))
562
+
563
+
564
+ _SVD_BUCKETS = (2, 4, 8, 16, 32, 64, 128, 256, 512, 1024)
565
+
566
+
567
+ def _bucket_sizes(max_bond: int) -> Tuple[int, ...]:
568
+ """Smallest-to-largest candidate SVD sizes for one MPSSimulator
569
+ instance's whole lifetime -- always ending at max_bond itself (added
570
+ if not already a power of 2 in _SVD_BUCKETS), so every real bond
571
+ dimension up to the hard cap has some bucket that fits it."""
572
+ buckets = tuple(b for b in _SVD_BUCKETS if b <= max_bond)
573
+ if not buckets or buckets[-1] != max_bond:
574
+ buckets = buckets + (max_bond,)
575
+ return buckets
576
+
577
+
578
+ def _build_mps_runner(n_qubits: int, max_bond: int, eps: float, jsd_budget: float):
579
+ """Factory: builds and returns a single @jax.jit-compiled closure that
580
+ runs an entire pre-compiled MPS circuit via jax.lax.scan -- same
581
+ factory pattern as chunk.py's _build_multi_chunk_runner, closing over
582
+ the static n_qubits/max_bond/eps/jsd_budget (fixed for one
583
+ MPSSimulator instance's whole lifetime, so the returned closure is
584
+ built once and cached on the instance, never rebuilt per call).
585
+
586
+ step's branch_1q returns a placeholder diag tuple (zeros) for 1-qubit
587
+ steps -- run_circuit_jit filters these out by g_id (>=20) before using
588
+ the per-step history, same as _bond_history/jsd_per_bond only ever
589
+ growing on 2-qubit gates in the eager path.
590
+
591
+ 2-qubit steps no longer always run the SVD at a fixed max_bond*2
592
+ size -- previously theta_mat.reshape(max_bond*2, 2*max_bond) was
593
+ ALWAYS square and ALWAYS a full economy SVD at that size, regardless
594
+ of the real/masked bond dimension (confirmed by direct measurement:
595
+ raising max_bond 64->128 measurably slowed every gate down even when
596
+ the real bond stayed ~4-8 on the same circuit). Instead, a small fixed
597
+ set of candidate sizes (_bucket_sizes) is dispatched via
598
+ jax.lax.switch -- same JIT-compatible mechanism _mps_1q_matrix/
599
+ _mps_2q_matrix already use for gate-ID dispatch, applied here to SVD
600
+ size instead. Verified correct across ~92 single-gate configurations
601
+ plus full multi-gate circuit fidelity (0.999999999998+) against this
602
+ module's own eager path, and 68.80x-73.96x faster on CPU / 2.74x on
603
+ GPU on a real N=50 TFIM Trotter circuit -- see Dense-Evolution-
604
+ Discovery's mps_bucketed_svd_optimization experiment for the full
605
+ validation, including a real bug found and fixed there before
606
+ promotion (see the bound formula's own comment below).
607
+
608
+ Bucket selection needs a provable (not heuristic) bound on how big B
609
+ must be, computed BEFORE the SVD from already-known real bond sizes
610
+ (real_chi, threaded through the scan carry alongside gammas/lambdas):
611
+ a 2-qubit gate acting at a cut can increase that cut's Schmidt rank by
612
+ at most a factor of d^2=4 (d=2, local physical dimension) -- theta_mat's
613
+ own matrix rank is bounded by min(chi_l_real*2, chi_r_real*2), a
614
+ mathematical fact about rank <= min(rows, cols), independent of the
615
+ PRE-gate middle bond. That output-side bound alone is not sufficient,
616
+ though: g1/g2/lam_m get sliced to size B BEFORE the einsum contracts
617
+ the middle bond away, so B must ALSO be at least as large as the
618
+ current middle bond (chi_m_real) and both outer bonds themselves, or
619
+ real (nonzero) Schmidt weight already on those bonds gets silently
620
+ dropped from the contraction -- not merely under-grown. Both
621
+ requirements combined: max(chi_l_real, chi_r_real, chi_m_real,
622
+ min(chi_l_real*2, chi_r_real*2))."""
623
+
624
+ buckets = _bucket_sizes(max_bond)
625
+ bucket_arr = jnp.array(buckets)
626
+
627
+ def step(carry, row):
628
+ gammas, lambdas, real_chi = carry
629
+ dtype = gammas.dtype
630
+
631
+ g_id = row[0].astype(jnp.int32)
632
+ q1 = row[1].astype(jnp.int32)
633
+ q2 = row[2].astype(jnp.int32)
634
+ param = row[3]
635
+ transpose_flag = row[4] > 0.5
636
+
637
+ def branch_1q(c):
638
+ gammas_, lambdas_, real_chi_ = c
639
+ gate_1q = _mps_1q_matrix(g_id, param, dtype)
640
+ new_g = jnp.einsum('ij,ljr->lir', gate_1q, gammas_[q1])
641
+ new_carry = (gammas_.at[q1].set(new_g), lambdas_, real_chi_)
642
+ real_dtype = _real_dtype_for(dtype)
643
+ diag = (jnp.asarray(0, dtype=jnp.int32), jnp.asarray(0.0, dtype=real_dtype),
644
+ jnp.asarray(0.0, dtype=real_dtype), jnp.asarray(0.0, dtype=real_dtype))
645
+ return new_carry, diag
646
+
647
+ def branch_2q(c):
648
+ # Traced unconditionally only when this branch is taken --
649
+ # jax.lax.cond (not jnp.where) skips the untaken branch's
650
+ # compute entirely, same choice compiler.py's do_2q makes for
651
+ # its own outer 1q/2q split: the SVD here is real work, worth
652
+ # skipping for the common case of a 1-qubit gate.
653
+ gammas_, lambdas_, real_chi_ = c
654
+ gate_2q = _mps_2q_matrix(g_id, param, dtype)
655
+ gate_2q = jnp.where(transpose_flag, jnp.transpose(gate_2q, (1, 0, 3, 2)), gate_2q)
656
+
657
+ chi_l_real = real_chi_[q1]
658
+ chi_m_real = real_chi_[q2] # current middle bond, BEFORE this gate
659
+ chi_r_real = real_chi_[q2 + 1]
660
+ input_min = jnp.maximum(jnp.maximum(chi_l_real, chi_r_real), chi_m_real)
661
+ output_bound = jnp.minimum(chi_l_real * 2, chi_r_real * 2)
662
+ bound = jnp.maximum(input_min, output_bound)
663
+ ge_mask = bucket_arr >= bound
664
+ bucket_idx = jnp.where(jnp.any(ge_mask), jnp.argmax(ge_mask), len(buckets) - 1)
665
+
666
+ def make_branch(B):
667
+ def branch_fn(_):
668
+ g1 = gammas_[q1][:B, :, :B]
669
+ g2 = gammas_[q2][:B, :, :B]
670
+ lam_l = lambdas_[q1][:B]
671
+ lam_m = lambdas_[q2][:B]
672
+ lam_r = lambdas_[q2 + 1][:B]
673
+ # Full Vidal update (prog.txt P0 fix): both outer
674
+ # Lambdas attached, same convention as the eager
675
+ # apply_gate_2q path -- see that method's docstring.
676
+ theta = jnp.einsum('l,lik,k,kjr,r->lijr', lam_l, g1, lam_m, g2, lam_r)
677
+ theta_new = jnp.einsum('abcd,ecdf->eabf', gate_2q, theta)
678
+ theta_mat = theta_new.reshape(B * 2, 2 * B)
679
+
680
+ U, S, Vh = jnp.linalg.svd(theta_mat, full_matrices=False)
681
+ chi_new, jsd_val = _vectorized_chi_search_jax(S, eps, jsd_budget, min(B, max_bond))
682
+ col_mask = jnp.arange(B) < chi_new
683
+
684
+ norm_full = jnp.sqrt(jnp.sum(S ** 2) + 1e-30)
685
+ S_norm_full = S / (norm_full + 1e-30)
686
+ trunc_err = jnp.sqrt(jnp.sum(jnp.where(jnp.arange(2 * B) >= chi_new, S_norm_full ** 2, 0.0)))
687
+
688
+ S_kept_masked = jnp.where(col_mask, S[:B], 0.0)
689
+ kept_norm = jnp.sqrt(jnp.sum(S_kept_masked ** 2) + 1e-30)
690
+ S_fixed = jnp.where(col_mask, S_kept_masked / (kept_norm + 1e-30), 0.0)
691
+
692
+ lam_l_inv = jnp.where(lam_l > eps, 1.0 / lam_l, 0.0)
693
+ lam_r_inv = jnp.where(lam_r > eps, 1.0 / lam_r, 0.0)
694
+
695
+ U_masked = jnp.where(col_mask[None, :], U[:, :B], 0.0)
696
+ Vh_masked = jnp.where(col_mask[:, None], Vh[:B, :], 0.0)
697
+
698
+ new_g1 = jnp.einsum('l,lir->lir', lam_l_inv, U_masked.reshape(B, 2, B))
699
+ new_g2 = jnp.einsum('ljr,r->ljr', Vh_masked.reshape(B, 2, B), lam_r_inv)
700
+
701
+ p_dist = S_fixed ** 2 # already unit-norm (see eager _svd_truncate)
702
+ ee = -jnp.sum(jnp.where(p_dist > 1e-20, p_dist * jnp.log2(jnp.where(p_dist > 1e-20, p_dist, 1.0)), 0.0))
703
+
704
+ real_dtype = _real_dtype_for(dtype)
705
+ return (_pad_gamma(new_g1, max_bond), _pad_gamma(new_g2, max_bond),
706
+ _pad_lambda(S_fixed, max_bond), chi_new.astype(jnp.int32),
707
+ jsd_val.astype(real_dtype), trunc_err.astype(real_dtype), ee.astype(real_dtype))
708
+
709
+ return branch_fn
710
+
711
+ branches = [make_branch(B) for B in buckets]
712
+ new_g1_p, new_g2_p, S_fixed_p, chi_new, jsd_val, trunc_err, ee = jax.lax.switch(
713
+ bucket_idx, branches, operand=None)
714
+
715
+ new_gammas = gammas_.at[q1].set(new_g1_p).at[q2].set(new_g2_p)
716
+ new_lambdas = lambdas_.at[q2].set(S_fixed_p)
717
+ new_real_chi = real_chi_.at[q2].set(chi_new)
718
+
719
+ new_carry = (new_gammas, new_lambdas, new_real_chi)
720
+ diag = (chi_new, jsd_val, trunc_err, ee)
721
+ return new_carry, diag
722
+
723
+ is_2q = g_id >= 20
724
+ new_carry, diag = jax.lax.cond(is_2q, branch_2q, branch_1q, carry)
725
+ return new_carry, diag
726
+
727
+ @jax.jit
728
+ def run(gammas, lambdas, real_chi, compiled_ops):
729
+ (final_gammas, final_lambdas, final_real_chi), diag = jax.lax.scan(
730
+ step, (gammas, lambdas, real_chi), compiled_ops)
731
+ return final_gammas, final_lambdas, final_real_chi, diag
732
+
733
+ return run
734
+
735
+
736
+ def _build_fused_mps_runner(n_qubits: int, max_bond: int, eps: float, jsd_budget: float):
737
+ """Factory for run_circuit_jit(ops, fuse_gates=True) -- same bucketed-
738
+ SVD dispatch as _build_mps_runner, except every step's gate matrix is
739
+ taken directly from a pre-fused matrix stream (built host-side by
740
+ _fuse_compiled_rows/_fused_entries_to_arrays from _compile_mps_ops's
741
+ own output) instead of being reconstructed inside the JIT via
742
+ _mps_1q_matrix/_mps_2q_matrix -- gate-ID dispatch is removed from the
743
+ traced program entirely, not just extended. Fusing consecutive
744
+ same-qubit-pair gates into one matrix means fewer, larger scan steps:
745
+ a real, measured ~2x GPU speedup on top of the bucketed dispatch
746
+ alone (see Dense-Evolution-Discovery's mps_gate_blocking_redesign_v2
747
+ experiment), on top of jax.lax.scan's own per-step GPU dispatch
748
+ overhead being the real bottleneck the bucketed dispatch alone
749
+ couldn't remove.
750
+
751
+ Returns the same (chi_new, jsd_val, trunc_err, entanglement_entropy)
752
+ diagnostic tuple _build_mps_runner does, one entry per FUSED step
753
+ rather than per original gate -- run_circuit_jit populates
754
+ self._bond_history/jsd_per_bond/truncation_errors/entanglement_entropy
755
+ from this at that coarser granularity when fuse_gates=True, an
756
+ explicit, documented trade-off (see MPSSimulator.run_circuit_jit's
757
+ own docstring), not a silent behavior change."""
758
+ buckets = _bucket_sizes(max_bond)
759
+ bucket_arr = jnp.array(buckets)
760
+
761
+ def step(carry, xs):
762
+ gammas, lambdas, real_chi = carry
763
+ dtype = gammas.dtype
764
+ is_2q, q1, q2, mat2q, mat1q = xs
765
+ q1 = q1.astype(jnp.int32)
766
+ q2 = q2.astype(jnp.int32)
767
+ real_dtype = _real_dtype_for(dtype)
768
+
769
+ def branch_1q(c):
770
+ gammas_, lambdas_, real_chi_ = c
771
+ new_g = jnp.einsum('ij,ljr->lir', mat1q, gammas_[q1])
772
+ new_carry = (gammas_.at[q1].set(new_g), lambdas_, real_chi_)
773
+ diag = (jnp.asarray(0, dtype=jnp.int32), jnp.asarray(0.0, dtype=real_dtype),
774
+ jnp.asarray(0.0, dtype=real_dtype), jnp.asarray(0.0, dtype=real_dtype))
775
+ return new_carry, diag
776
+
777
+ def branch_2q(c):
778
+ gammas_, lambdas_, real_chi_ = c
779
+ gate_2q = mat2q
780
+
781
+ chi_l_real = real_chi_[q1]
782
+ chi_m_real = real_chi_[q2]
783
+ chi_r_real = real_chi_[q2 + 1]
784
+ input_min = jnp.maximum(jnp.maximum(chi_l_real, chi_r_real), chi_m_real)
785
+ output_bound = jnp.minimum(chi_l_real * 2, chi_r_real * 2)
786
+ bound = jnp.maximum(input_min, output_bound)
787
+ ge_mask = bucket_arr >= bound
788
+ bucket_idx = jnp.where(jnp.any(ge_mask), jnp.argmax(ge_mask), len(buckets) - 1)
789
+
790
+ def make_branch(B):
791
+ def branch_fn(_):
792
+ g1 = gammas_[q1][:B, :, :B]
793
+ g2 = gammas_[q2][:B, :, :B]
794
+ lam_l = lambdas_[q1][:B]
795
+ lam_m = lambdas_[q2][:B]
796
+ lam_r = lambdas_[q2 + 1][:B]
797
+ theta = jnp.einsum('l,lik,k,kjr,r->lijr', lam_l, g1, lam_m, g2, lam_r)
798
+ theta_new = jnp.einsum('abcd,ecdf->eabf', gate_2q, theta)
799
+ theta_mat = theta_new.reshape(B * 2, 2 * B)
800
+
801
+ U, S, Vh = jnp.linalg.svd(theta_mat, full_matrices=False)
802
+ chi_new, jsd_val = _vectorized_chi_search_jax(S, eps, jsd_budget, min(B, max_bond))
803
+ col_mask = jnp.arange(B) < chi_new
804
+
805
+ norm_full = jnp.sqrt(jnp.sum(S ** 2) + 1e-30)
806
+ S_norm_full = S / (norm_full + 1e-30)
807
+ trunc_err = jnp.sqrt(jnp.sum(jnp.where(jnp.arange(2 * B) >= chi_new, S_norm_full ** 2, 0.0)))
808
+
809
+ S_kept_masked = jnp.where(col_mask, S[:B], 0.0)
810
+ kept_norm = jnp.sqrt(jnp.sum(S_kept_masked ** 2) + 1e-30)
811
+ S_fixed = jnp.where(col_mask, S_kept_masked / (kept_norm + 1e-30), 0.0)
812
+
813
+ lam_l_inv = jnp.where(lam_l > eps, 1.0 / lam_l, 0.0)
814
+ lam_r_inv = jnp.where(lam_r > eps, 1.0 / lam_r, 0.0)
815
+
816
+ U_masked = jnp.where(col_mask[None, :], U[:, :B], 0.0)
817
+ Vh_masked = jnp.where(col_mask[:, None], Vh[:B, :], 0.0)
818
+
819
+ new_g1 = jnp.einsum('l,lir->lir', lam_l_inv, U_masked.reshape(B, 2, B))
820
+ new_g2 = jnp.einsum('ljr,r->ljr', Vh_masked.reshape(B, 2, B), lam_r_inv)
821
+
822
+ p_dist = S_fixed ** 2
823
+ ee = -jnp.sum(jnp.where(p_dist > 1e-20, p_dist * jnp.log2(jnp.where(p_dist > 1e-20, p_dist, 1.0)), 0.0))
824
+
825
+ return (_pad_gamma(new_g1, max_bond), _pad_gamma(new_g2, max_bond),
826
+ _pad_lambda(S_fixed, max_bond), chi_new.astype(jnp.int32),
827
+ jsd_val.astype(real_dtype), trunc_err.astype(real_dtype), ee.astype(real_dtype))
828
+
829
+ return branch_fn
830
+
831
+ branches = [make_branch(B) for B in buckets]
832
+ new_g1_p, new_g2_p, S_fixed_p, chi_new, jsd_val, trunc_err, ee = jax.lax.switch(
833
+ bucket_idx, branches, operand=None)
834
+
835
+ new_gammas = gammas_.at[q1].set(new_g1_p).at[q2].set(new_g2_p)
836
+ new_lambdas = lambdas_.at[q2].set(S_fixed_p)
837
+ new_real_chi = real_chi_.at[q2].set(chi_new)
838
+
839
+ new_carry = (new_gammas, new_lambdas, new_real_chi)
840
+ diag = (chi_new, jsd_val, trunc_err, ee)
841
+ return new_carry, diag
842
+
843
+ new_carry, diag = jax.lax.cond(is_2q, branch_2q, branch_1q, carry)
844
+ return new_carry, diag
845
+
846
+ @jax.jit
847
+ def run(gammas, lambdas, real_chi, xs):
848
+ (final_gammas, final_lambdas, final_real_chi), diag = jax.lax.scan(
849
+ step, (gammas, lambdas, real_chi), xs)
850
+ return final_gammas, final_lambdas, final_real_chi, diag
851
+
852
+ return run
853
+
854
+
855
+ class MPSSimulator:
856
+ """
857
+ Matrix Product State simulator with adaptive SVD-truncated bond
858
+ dimension (JSD-budget driven), no lossy post-truncation quantization.
859
+ JAX-backed core (einsum, SVD).
860
+
861
+ Parameters
862
+ ----------
863
+ n_qubits : int
864
+ max_bond : int -- hard cap on bond dimension chi
865
+ svd_cutoff : float or None -- singular values below this are dropped
866
+ outright. None (default) resolves to a value
867
+ appropriate for the dtype this instance actually
868
+ runs at: ~1e-12 for complex128, ~1e-6 (roughly
869
+ 10x float32's own machine epsilon) for complex64
870
+ -- a fixed 1e-12 in complex64 sits below that
871
+ dtype's noise floor, so numerical noise gets
872
+ counted as real Schmidt weight and chi never
873
+ shrinks below max_bond regardless of jsd_budget.
874
+ An explicitly passed value always wins verbatim,
875
+ never rescaled.
876
+ jsd_budget : float -- max tolerated Jensen-Shannon distance between the
877
+ full and truncated singular-value distributions
878
+ at each cut; chi is grown by 1 until satisfied
879
+ or max_bond is hit.
880
+ use_float32 : bool or None -- None (default) follows the process-wide
881
+ jax_enable_x64 flag, same convention as this
882
+ module always used (see module docstring).
883
+ True forces complex64 (and the complex64-
884
+ appropriate svd_cutoff default) even if x64 is
885
+ enabled. False requests complex128 -- since
886
+ complex128 arrays don't exist in JAX at all
887
+ unless the process-wide flag is on, this calls
888
+ the same lazy ensure_x64() DenseSVSimulator
889
+ uses, mirroring its own use_float32=False
890
+ handling. That call is a no-op if precision was
891
+ already pinned via set_precision(), same
892
+ deference DenseSVSimulator gives it; the flag is
893
+ re-read afterward rather than assumed, so
894
+ dtype/eps stay consistent with whatever
895
+ precision is really active even in that
896
+ pinned-False edge case.
897
+ """
898
+
899
+ def __init__(
900
+ self,
901
+ n_qubits: int,
902
+ max_bond: int = 64,
903
+ svd_cutoff: Optional[float] = None,
904
+ jsd_budget: float = 1e-5,
905
+ use_float32: Optional[bool] = None,
906
+ ):
907
+ self.n = n_qubits
908
+ self.chi = max_bond
909
+ if use_float32 is None:
910
+ x64_active = jax.config.jax_enable_x64
911
+ elif use_float32:
912
+ x64_active = False
913
+ else:
914
+ ensure_x64()
915
+ x64_active = jax.config.jax_enable_x64
916
+ dtype = jnp.complex128 if x64_active else jnp.complex64
917
+ self.eps = svd_cutoff if svd_cutoff is not None else (1e-12 if x64_active else 1e-6)
918
+ self.jsd_budget = jsd_budget
919
+
920
+ self.gammas: List[jnp.ndarray] = []
921
+ self.lambdas: List[jnp.ndarray] = [jnp.ones(1)] * (n_qubits + 1)
922
+ # Real (non-max_bond-padded) bond dimension at every cut, kept in
923
+ # sync by BOTH the eager path (apply_gate_2q, below) and
924
+ # run_circuit_jit -- needed by the bucketed-SVD dispatch in
925
+ # _build_mps_runner to pick a provably-sufficient bucket size
926
+ # without ever inferring it from zero-counting (see that
927
+ # function's own docstring for why that would be unreliable).
928
+ self._real_chi: np.ndarray = np.ones(n_qubits + 1, dtype=np.int64)
929
+
930
+ self.truncation_errors: List[float] = []
931
+ self.jsd_per_bond: List[float] = []
932
+ self.entanglement_entropy = np.zeros(max(n_qubits - 1, 0))
933
+ self._bond_history: List[int] = []
934
+ # Counts truncations where max_bond was hit before jsd_budget could
935
+ # be satisfied -- the while loop below exits silently in that case,
936
+ # and avg_JSD (a mean over all steps) can look deceptively low even
937
+ # when the final contracted state is badly wrong (verified: TVD
938
+ # ~0.97 against DenseSVSimulator on an 8-qubit/15-layer entangling
939
+ # circuit with max_bond=2, while avg_JSD read 0.0534).
940
+ self.budget_violations: int = 0
941
+
942
+ # Cached compiled closure for run_circuit_jit -- built lazily on
943
+ # first use (self.n/self.chi/self.eps/self.jsd_budget are fixed
944
+ # for this instance's lifetime), never rebuilt per call. Same
945
+ # caching pattern as Chunk.__init__'s self._multi_chunk_runner.
946
+ self._mps_runner = None
947
+ self._fused_mps_runner = None
948
+
949
+ for _ in range(n_qubits):
950
+ g = jnp.zeros((1, 2, 1), dtype=dtype)
951
+ g = g.at[0, 0, 0].set(1.0)
952
+ self.gammas.append(g)
953
+
954
+ # gate 1q
955
+ def apply_gate_1q(self, gate: jnp.ndarray, qubit: int) -> None:
956
+ """O(chi^2) -- updates only Gamma[qubit]."""
957
+ gate = jnp.asarray(gate)
958
+ self.gammas[qubit] = jnp.einsum("ij,ljr->lir", gate, self.gammas[qubit])
959
+
960
+ # core: plain adaptive SVD truncation (no quantization)
961
+ def _svd_truncate(
962
+ self, theta_mat: jnp.ndarray
963
+ ) -> Tuple[jnp.ndarray, jnp.ndarray, jnp.ndarray, float, float]:
964
+ """SVD + adaptive chi search, plus the two corrections from prog.txt
965
+ P0: trunc_err/entanglement_entropy need the true (normalized)
966
+ Schmidt spectrum, and the kept singular values are renormalized to
967
+ unit norm after truncation (same convention Vidal's iTEBD update
968
+ uses -- the state's overall norm is the product of every bond's
969
+ Lambda norm, each already implicitly 1 by induction, so explicitly
970
+ restoring that here keeps the whole chain normalized without ever
971
+ relying on a global re-norm in contract_to_statevector)."""
972
+ U, S, Vh = jnp.linalg.svd(theta_mat, full_matrices=False)
973
+
974
+ # Was a Python while loop incrementing chi_new one at a time,
975
+ # host-syncing (float()/int() casts) every iteration -- the
976
+ # eager path never used the vectorized search already written
977
+ # for the JIT path (_vectorized_chi_search / ..._jax), despite
978
+ # that search being verified to replicate the while loop's
979
+ # chi_new/jsd_val exactly (0 mismatches across 171 real calls,
980
+ # including this budget-violation fallback -- see that
981
+ # function's own docstring).
982
+ chi_new, jsd_val = _vectorized_chi_search(S, self.eps, self.jsd_budget, self.chi)
983
+ max_possible = min(len(S), self.chi)
984
+
985
+ if jsd_val > self.jsd_budget:
986
+ # Loop exited because chi_new hit max_possible (== max_bond, in
987
+ # the common case where the bond isn't already capped by the
988
+ # SVD's own rank), not because jsd_budget was satisfied.
989
+ if self.budget_violations == 0:
990
+ warnings.warn(
991
+ f"MPSSimulator: bond dimension capped at max_bond={self.chi}, "
992
+ f"jsd_budget={self.jsd_budget:.1e} not honored "
993
+ f"(jsd={jsd_val:.2e}) -- results may be unreliable, "
994
+ f"consider raising max_bond.",
995
+ UserWarning,
996
+ stacklevel=2,
997
+ )
998
+ self.budget_violations += 1
999
+
1000
+ norm_full = float(jnp.sqrt(jnp.sum(S ** 2) + 1e-30))
1001
+ S_norm_full = S / (norm_full + 1e-30)
1002
+ trunc_err = float(jnp.sqrt(jnp.sum(S_norm_full[chi_new:] ** 2))) if len(S) > chi_new else 0.0
1003
+ self.truncation_errors.append(trunc_err)
1004
+
1005
+ S_kept = S[:chi_new]
1006
+ kept_norm = float(jnp.sqrt(jnp.sum(S_kept ** 2) + 1e-30))
1007
+ S_kept_renorm = S_kept / (kept_norm + 1e-30)
1008
+
1009
+ return U[:, :chi_new], S_kept_renorm, Vh[:chi_new, :], trunc_err, jsd_val
1010
+
1011
+ # gate 2q adjacent
1012
+ def apply_gate_2q(self, gate_2q: jnp.ndarray, q1: int, q2: int) -> None:
1013
+ """2-qubit gate with adaptive SVD truncation. O(chi^3).
1014
+
1015
+ Vidal's full two-site update (prog.txt P0 fix): theta is built from
1016
+ BOTH outer Lambdas (Lambda[q1], Lambda[q2+1]) as well as the middle
1017
+ one, not just the middle one -- so its singular values are the true
1018
+ global Schmidt coefficients at this cut, not an artifact of the
1019
+ local 2-site reduced state. New Gamma tensors are recovered by
1020
+ dividing the outer Lambdas back out (regularized: entries at or
1021
+ below svd_cutoff map to a zero inverse instead of blowing up --
1022
+ exactly the padded/zero entries in the JIT path's fixed-size
1023
+ arrays, and the trivial size-1 boundary Lambda everywhere else)."""
1024
+ gate_2q = jnp.asarray(gate_2q)
1025
+ if abs(q1 - q2) != 1:
1026
+ self._apply_nonlocal_2q(gate_2q, q1, q2)
1027
+ return
1028
+ if q1 > q2:
1029
+ q1, q2 = q2, q1
1030
+ gate_2q = jnp.transpose(gate_2q, (1, 0, 3, 2))
1031
+
1032
+ lam_l = self.lambdas[q1]
1033
+ g1 = self.gammas[q1]
1034
+ lam_m = self.lambdas[q2]
1035
+ g2 = self.gammas[q2]
1036
+ lam_r = self.lambdas[q2 + 1]
1037
+
1038
+ theta = jnp.einsum("l,lik,k,kjr,r->lijr", lam_l, g1, lam_m, g2, lam_r)
1039
+ chiL, d1, d2, chiR = theta.shape
1040
+
1041
+ theta_new = jnp.einsum("abcd,ecdf->eabf", gate_2q, theta)
1042
+ theta_mat = theta_new.reshape(chiL * d1, d2 * chiR)
1043
+
1044
+ U_t, S_t, Vh_t, trunc_err, jsd_val = self._svd_truncate(theta_mat)
1045
+ chi_new = len(S_t)
1046
+
1047
+ lam_l_inv = jnp.where(lam_l > self.eps, 1.0 / lam_l, 0.0)
1048
+ lam_r_inv = jnp.where(lam_r > self.eps, 1.0 / lam_r, 0.0)
1049
+
1050
+ new_g1 = jnp.einsum("l,lir->lir", lam_l_inv, U_t.reshape(chiL, d1, chi_new))
1051
+ new_g2 = jnp.einsum("ljr,r->ljr", Vh_t.reshape(chi_new, d2, chiR), lam_r_inv)
1052
+
1053
+ p_dist = S_t ** 2 # already unit-norm (see _svd_truncate)
1054
+ mask = p_dist > 1e-20
1055
+ ee = float(-jnp.sum(jnp.where(mask, p_dist * jnp.log2(jnp.where(mask, p_dist, 1.0)), 0.0)))
1056
+ if q1 < len(self.entanglement_entropy):
1057
+ self.entanglement_entropy[q1] = ee
1058
+
1059
+ self.lambdas[q2] = S_t
1060
+ self.gammas[q1] = new_g1
1061
+ self.gammas[q2] = new_g2
1062
+ self._real_chi[q2] = chi_new
1063
+
1064
+ self._bond_history.append(chi_new)
1065
+ self.jsd_per_bond.append(jsd_val)
1066
+
1067
+ def _apply_nonlocal_2q(self, gate_2q: jnp.ndarray, q1: int, q2: int) -> None:
1068
+ """Non-adjacent 2-qubit gate via a SWAP chain down to adjacent.
1069
+
1070
+ The SWAP chain always ends up applying the gate at (target, target+1)
1071
+ with target = min(q1, q2) first -- so without normalizing here, a
1072
+ caller passing q1 > q2 (e.g. apply_cx(ctrl=3, tgt=1)) would silently
1073
+ have its control/target roles swapped for asymmetric gates like CNOT.
1074
+ Same fix as apply_gate_2q's adjacent-qubit branch just above, applied
1075
+ before the swap chain runs so it always sees q1 < q2."""
1076
+ if q1 > q2:
1077
+ q1, q2 = q2, q1
1078
+ gate_2q = jnp.transpose(gate_2q, (1, 0, 3, 2))
1079
+ swap = jnp.array(
1080
+ [[1, 0, 0, 0], [0, 0, 1, 0], [0, 1, 0, 0], [0, 0, 0, 1]], dtype=complex
1081
+ ).reshape(2, 2, 2, 2)
1082
+ target = min(q1, q2)
1083
+ for q in range(max(q1, q2) - 1, target, -1):
1084
+ self.apply_gate_2q(swap, q, q + 1)
1085
+ self.apply_gate_2q(gate_2q, target, target + 1)
1086
+ for q in range(target + 1, max(q1, q2)):
1087
+ self.apply_gate_2q(swap, q, q + 1)
1088
+
1089
+ # gate shortcuts
1090
+ def apply_cx(self, ctrl: int, tgt: int) -> None:
1091
+ cx = jnp.array(
1092
+ [[1, 0, 0, 0], [0, 1, 0, 0], [0, 0, 0, 1], [0, 0, 1, 0]], dtype=complex
1093
+ ).reshape(2, 2, 2, 2)
1094
+ self.apply_gate_2q(cx, ctrl, tgt)
1095
+
1096
+ def apply_cz(self, ctrl: int, tgt: int) -> None:
1097
+ cz = jnp.diag(jnp.array([1, 1, 1, -1])).astype(complex).reshape(2, 2, 2, 2)
1098
+ self.apply_gate_2q(cz, ctrl, tgt)
1099
+
1100
+ def apply_swap(self, q1: int, q2: int) -> None:
1101
+ sw = jnp.array(
1102
+ [[1, 0, 0, 0], [0, 0, 1, 0], [0, 1, 0, 0], [0, 0, 0, 1]], dtype=complex
1103
+ ).reshape(2, 2, 2, 2)
1104
+ self.apply_gate_2q(sw, q1, q2)
1105
+
1106
+ def apply_ccx(self, c1: int, c2: int, tgt: int) -> None:
1107
+ """Toffoli via standard T-gate decomposition (all 1q/2q gates)."""
1108
+ inv2 = 1.0 / np.sqrt(2.0)
1109
+ h = inv2 * jnp.array([[1, 1], [1, -1]], dtype=complex)
1110
+ t = jnp.array([[1, 0], [0, jnp.exp(1j * jnp.pi / 4)]], dtype=complex)
1111
+ tdg = jnp.array([[1, 0], [0, jnp.exp(-1j * jnp.pi / 4)]], dtype=complex)
1112
+
1113
+ self.apply_gate_1q(h, tgt)
1114
+ self.apply_cx(c2, tgt)
1115
+ self.apply_gate_1q(tdg, tgt)
1116
+ self.apply_cx(c1, tgt)
1117
+ self.apply_gate_1q(t, tgt)
1118
+ self.apply_cx(c2, tgt)
1119
+ self.apply_gate_1q(tdg, tgt)
1120
+ self.apply_cx(c1, tgt)
1121
+ self.apply_gate_1q(t, c2)
1122
+ self.apply_gate_1q(t, tgt)
1123
+ self.apply_gate_1q(h, tgt)
1124
+ self.apply_cx(c1, c2)
1125
+ self.apply_gate_1q(t, c1)
1126
+ self.apply_gate_1q(tdg, c2)
1127
+ self.apply_cx(c1, c2)
1128
+
1129
+ # contraction to a full statevector -- O(2**n), n <= ~24 only
1130
+ def contract_to_statevector(self) -> jnp.ndarray:
1131
+ if self.n > 24:
1132
+ raise MemoryError(
1133
+ f"contract_to_statevector at {self.n} qubits would need "
1134
+ f"~{2**self.n * 16 / 1e9:.1f} GB -- use "
1135
+ f"get_probabilities_sampled or get_top_k_probable_states "
1136
+ f"instead for n > 24."
1137
+ )
1138
+ # Index [0] rather than .squeeze(axis=0)/.squeeze(axis=-1): the
1139
+ # boundary axis has size 1 in the eager (unpadded) representation
1140
+ # (squeeze and [0]-indexing are identical there), but size
1141
+ # max_bond after run_circuit_jit's fixed-shape padding -- the real
1142
+ # boundary content is always at index 0 by construction either
1143
+ # way (see _pad_gamma), so [0]-indexing is the strict
1144
+ # generalization that works for both, not a behavior change for
1145
+ # the pre-existing eager path.
1146
+ result = self.gammas[0][0, :, :]
1147
+ for i in range(1, self.n):
1148
+ lam = self.lambdas[i]
1149
+ g = self.gammas[i]
1150
+ lg = jnp.einsum("k,kir->kir", lam, g)
1151
+ result = jnp.tensordot(result, lg, axes=([-1], [0]))
1152
+ result = result.reshape(-1, result.shape[-1])
1153
+ sv = result[..., 0]
1154
+ norm = jnp.linalg.norm(sv)
1155
+ return sv / (norm + 1e-15)
1156
+
1157
+ # sequential sampling -- O(chi^2) per qubit per sample, never
1158
+ # materializes a (2**n,) array. The only way to get results for
1159
+ # n_qubits beyond ~24. Sampling decisions themselves stay on host
1160
+ # (np.random.Generator): each step needs a concrete probability to
1161
+ # branch on, not a traced value, so this loop isn't a jax.jit
1162
+ # candidate as a whole regardless of backend.
1163
+ def _sample_bitstring(self, rng: np.random.Generator) -> List[int]:
1164
+ bits = []
1165
+ state = jnp.ones(1, dtype=self.gammas[0].dtype)
1166
+ for i in range(self.n):
1167
+ g = self.gammas[i]
1168
+ lam = self.lambdas[i + 1] if i < self.n - 1 else jnp.ones(g.shape[2])
1169
+ p0v = jnp.einsum("l,lr->r", state, g[:, 0, :])
1170
+ p1v = jnp.einsum("l,lr->r", state, g[:, 1, :])
1171
+ p0l = p0v * lam
1172
+ p1l = p1v * lam
1173
+ p0 = float(jnp.real(jnp.dot(p0l, jnp.conj(p0l))))
1174
+ p1 = float(jnp.real(jnp.dot(p1l, jnp.conj(p1l))))
1175
+ norm = p0 + p1 + 1e-15
1176
+ bit = 0 if rng.random() < p0 / norm else 1
1177
+ bits.append(bit)
1178
+ state = p0l if bit == 0 else p1l
1179
+ state = state / (jnp.linalg.norm(state) + 1e-15)
1180
+ return bits
1181
+
1182
+ def get_probabilities_sampled(
1183
+ self, n_samples: int = 100_000, seed: Optional[int] = None
1184
+ ) -> dict:
1185
+ """Returns a {bitstring: empirical_probability} dict from n_samples
1186
+ sequential draws -- the only entry point safe for n_qubits > 24."""
1187
+ from collections import Counter
1188
+
1189
+ rng = np.random.default_rng(seed)
1190
+ counts: Counter = Counter()
1191
+ for _ in range(n_samples):
1192
+ bits = self._sample_bitstring(rng)
1193
+ counts["".join(map(str, bits))] += 1
1194
+ return {bitstr: c / n_samples for bitstr, c in counts.items()}
1195
+
1196
+ # ──────────────────────────────────────────────
1197
+ # approximate top-k extraction -- see module docstring point 2.
1198
+ # ──────────────────────────────────────────────
1199
+ def get_top_k_probable_states(self, k: int = 128) -> Tuple[np.ndarray, np.ndarray]:
1200
+ """Greedy beam search (beam width k) for approximately-most-probable
1201
+ basis states, without ever contracting to a full statevector.
1202
+
1203
+ Returns (indices, probabilities): indices are computational-basis
1204
+ integers, probabilities are exact for the states found (not
1205
+ approximated), sorted descending. Recall of the TRUE top states
1206
+ improves with k but is not guaranteed for any fixed k -- see the
1207
+ module docstring."""
1208
+ paths: List[Tuple[int, jnp.ndarray]] = [(0, jnp.array([1.0 + 0.0j]))]
1209
+ for i in range(self.n):
1210
+ candidates = []
1211
+ gamma = self.gammas[i]
1212
+ lam = self.lambdas[i + 1] if (i + 1) < len(self.lambdas) else jnp.ones(gamma.shape[2])
1213
+ for idx_p, vec_p in paths:
1214
+ for bit in (0, 1):
1215
+ new_vec = jnp.einsum("l,lr->r", vec_p, gamma[:, bit, :]) * lam
1216
+ weight = float(jnp.sum(jnp.abs(new_vec) ** 2))
1217
+ candidates.append(((idx_p << 1) | bit, new_vec, weight))
1218
+ # heapq.nlargest instead of a full sort-then-slice (prog.txt
1219
+ # point 5e): only the top k by weight are ever used below, and
1220
+ # this avoids materializing/sorting the full candidates list
1221
+ # when len(candidates) >> k. Final probability order is
1222
+ # re-derived from scratch at the end of this function anyway
1223
+ # (`order = np.argsort(-probabilities)`), so which of two
1224
+ # equal-weight candidates heapq happens to prefer over sort's
1225
+ # stable order has no effect on the result.
1226
+ paths = [(idx, vec) for idx, vec, _ in heapq.nlargest(k, candidates, key=lambda c: c[2])]
1227
+
1228
+ indices = np.array([p[0] for p in paths])
1229
+ amplitudes = np.array([
1230
+ complex(vec[0]) if len(vec) == 1 else complex(jnp.sum(vec))
1231
+ for _, vec in paths
1232
+ ])
1233
+ probabilities = np.abs(amplitudes) ** 2
1234
+ order = np.argsort(-probabilities)
1235
+ return indices[order], probabilities[order]
1236
+
1237
+ # metrics
1238
+ def max_bond_used(self) -> int:
1239
+ return max(self._bond_history) if self._bond_history else 1
1240
+
1241
+ def total_truncation_error(self) -> float:
1242
+ if not self.truncation_errors:
1243
+ return 0.0
1244
+ return float(np.sqrt(np.sum(np.array(self.truncation_errors) ** 2)))
1245
+
1246
+ def avg_jsd(self) -> float:
1247
+ return float(np.mean(self.jsd_per_bond)) if self.jsd_per_bond else 0.0
1248
+
1249
+ def memory_bytes(self) -> int:
1250
+ bytes_gammas = sum(g.size * g.dtype.itemsize for g in self.gammas)
1251
+ bytes_lambdas = sum(l.size * l.dtype.itemsize for l in self.lambdas)
1252
+ return int(bytes_gammas + bytes_lambdas)
1253
+
1254
+ def memory_mb(self) -> float:
1255
+ return self.memory_bytes() / (1024 * 1024)
1256
+
1257
+ def summary(self) -> str:
1258
+ ee_max = self.entanglement_entropy.max() if len(self.entanglement_entropy) else 0.0
1259
+ return (
1260
+ f"MPSSimulator | n={self.n} | chi_max={self.chi} | "
1261
+ f"chi_used={self.max_bond_used()} | mem={self.memory_mb():.3f}MB | "
1262
+ f"trunc_err={self.total_truncation_error():.2e} | "
1263
+ f"avg_JSD={self.avg_jsd():.4f} | EE_max={ee_max:.3f}b | "
1264
+ f"budget_violations={self.budget_violations}"
1265
+ )
1266
+
1267
+ # ── JIT-fused whole-circuit execution ─────────────────────────────
1268
+ def _record_diag_bookkeeping(self, diag, q1_ids: np.ndarray, is_2q_mask: np.ndarray) -> None:
1269
+ """Shared by both run_circuit_jit paths: jax.lax.scan's stacked
1270
+ per-step diagnostics (diag) replace the eager path's Python
1271
+ list.append()s inside the loop -- same final content, populated
1272
+ differently. Only 2-qubit steps count (is_2q_mask), same as
1273
+ _bond_history/jsd_per_bond only ever growing on 2-qubit gates in
1274
+ the eager path."""
1275
+ chi_history, jsd_history, trunc_err_history, entropy_history = (
1276
+ np.asarray(diag[0]), np.asarray(diag[1]), np.asarray(diag[2]), np.asarray(diag[3]))
1277
+ for i in np.nonzero(is_2q_mask)[0]:
1278
+ chi_new = int(chi_history[i])
1279
+ jsd_val = float(jsd_history[i])
1280
+ self._bond_history.append(chi_new)
1281
+ self.jsd_per_bond.append(jsd_val)
1282
+ self.truncation_errors.append(float(trunc_err_history[i]))
1283
+ q1 = q1_ids[i]
1284
+ if q1 < len(self.entanglement_entropy):
1285
+ self.entanglement_entropy[q1] = float(entropy_history[i])
1286
+ if jsd_val > self.jsd_budget:
1287
+ if self.budget_violations == 0:
1288
+ warnings.warn(
1289
+ f"MPSSimulator: bond dimension capped at max_bond={self.chi}, "
1290
+ f"jsd_budget={self.jsd_budget:.1e} not honored "
1291
+ f"(jsd={jsd_val:.2e}) -- results may be unreliable, "
1292
+ f"consider raising max_bond.",
1293
+ UserWarning,
1294
+ stacklevel=2,
1295
+ )
1296
+ self.budget_violations += 1
1297
+
1298
+ def run_circuit_jit(self, ops: List, fuse_gates: bool = False) -> None:
1299
+ """Runs an entire circuit through a single jax.lax.scan-fused,
1300
+ @jax.jit-compiled kernel instead of one eager Python call per gate
1301
+ -- the eager path (apply_gate_1q/apply_gate_2q/_apply_nonlocal_2q,
1302
+ all still available and unchanged) has zero @jax.jit anywhere and
1303
+ pays a host-device sync on every 2-qubit gate's bond-dimension
1304
+ search; measured 88.9s vs Qiskit Aer's 0.64s on a 60-qubit stress
1305
+ circuit -- see README changelog for the real before/after number
1306
+ this method produces on that same circuit.
1307
+
1308
+ Trade-off, explicit and intentional (not hidden): every gamma/
1309
+ lambda is kept at a fixed max_bond-padded size for the rest of
1310
+ this instance's lifetime after this call. Structurally correct
1311
+ either way (zero-padding is mathematically transparent to every
1312
+ other method here -- contract_to_statevector, get_top_k_probable_
1313
+ states, etc. all still work correctly on the padded arrays,
1314
+ verified), just not memory-minimal for genuinely low-entanglement
1315
+ circuits, which is this module's whole point for very large qubit
1316
+ counts. Use the eager methods directly instead when memory, not
1317
+ speed, is the priority -- this is an addition, not a replacement.
1318
+
1319
+ ops: same convention as DenseSVSimulator.run_circuit_jit_beast_mode
1320
+ -- list of (name, *args) tuples/lists. Unlike that method, SWAP is
1321
+ never decomposed into 3xCX (kept as one real gate, see
1322
+ _compile_mps_ops's docstring for why that matters here).
1323
+
1324
+ fuse_gates: opt-in, default False. When True, consecutive gates
1325
+ acting on the same (or a growing) qubit pair are fused into one
1326
+ matrix on the host before compiling (exact -- matrix
1327
+ multiplication, no approximation), cutting the number of scan
1328
+ steps and measurably faster on GPU (~2x on top of the bucketed
1329
+ SVD dispatch alone, see Dense-Evolution-Discovery's
1330
+ mps_gate_blocking_redesign_v2 experiment for the full validation,
1331
+ including verification against non-adjacent-gate and CCX
1332
+ circuits). The trade-off: self._bond_history/jsd_per_bond/
1333
+ truncation_errors/entanglement_entropy get one entry per FUSED
1334
+ step instead of per original gate -- real diagnostics, just
1335
+ coarser-grained. Defaults to False so existing behavior and
1336
+ per-gate bookkeeping granularity are unchanged unless requested.
1337
+ """
1338
+ dtype = self.gammas[0].dtype
1339
+ lambda_dtype = self.lambdas[0].dtype
1340
+
1341
+ if fuse_gates:
1342
+ compiled_rows = _compile_mps_ops(ops, self.n)
1343
+ fused = _fuse_compiled_rows(compiled_rows, dtype) if compiled_rows else []
1344
+
1345
+ if self._fused_mps_runner is None:
1346
+ self._fused_mps_runner = _build_fused_mps_runner(self.n, self.chi, self.eps, self.jsd_budget)
1347
+
1348
+ gammas_padded = _pad_all_gammas(tuple(self.gammas), self.chi, dtype)
1349
+ lambdas_padded = _pad_all_lambdas(tuple(self.lambdas), self.chi, lambda_dtype)
1350
+ real_chi_initial = jnp.asarray(self._real_chi, dtype=jnp.int32)
1351
+
1352
+ if fused:
1353
+ xs = _fused_entries_to_arrays(fused, dtype)
1354
+ final_gammas, final_lambdas, final_real_chi, diag = self._fused_mps_runner(
1355
+ gammas_padded, lambdas_padded, real_chi_initial, xs)
1356
+
1357
+ self.gammas = [final_gammas[i] for i in range(self.n)]
1358
+ self.lambdas = [final_lambdas[i] for i in range(self.n + 1)]
1359
+ self._real_chi = np.asarray(final_real_chi)
1360
+
1361
+ q1_ids = np.asarray([entry[1] for entry in fused])
1362
+ is_2q_mask = np.asarray([entry[0] == '2q' for entry in fused])
1363
+ self._record_diag_bookkeeping(diag, q1_ids, is_2q_mask)
1364
+ return
1365
+
1366
+ compiled_rows = _compile_mps_ops(ops, self.n)
1367
+ ops_dtype = _real_dtype_for(dtype)
1368
+
1369
+ if compiled_rows:
1370
+ ops_array = jnp.array(compiled_rows, dtype=ops_dtype)
1371
+ else:
1372
+ ops_array = jnp.zeros((0, 5), dtype=ops_dtype)
1373
+
1374
+ if self._mps_runner is None:
1375
+ self._mps_runner = _build_mps_runner(self.n, self.chi, self.eps, self.jsd_budget)
1376
+
1377
+ gammas_padded = jnp.stack([_pad_gamma(g, self.chi).astype(dtype) for g in self.gammas])
1378
+ lambdas_padded = jnp.stack([_pad_lambda(l, self.chi).astype(lambda_dtype) for l in self.lambdas])
1379
+ real_chi_initial = jnp.asarray(self._real_chi, dtype=jnp.int32)
1380
+
1381
+ final_gammas, final_lambdas, final_real_chi, diag = self._mps_runner(
1382
+ gammas_padded, lambdas_padded, real_chi_initial, ops_array)
1383
+
1384
+ self.gammas = [final_gammas[i] for i in range(self.n)]
1385
+ self.lambdas = [final_lambdas[i] for i in range(self.n + 1)]
1386
+ self._real_chi = np.asarray(final_real_chi)
1387
+
1388
+ if compiled_rows:
1389
+ g_ids = np.asarray([row[0] for row in compiled_rows])
1390
+ q1_ids = np.asarray([int(row[1]) for row in compiled_rows])
1391
+ is_2q_mask = g_ids >= 20
1392
+ self._record_diag_bookkeeping(diag, q1_ids, is_2q_mask)
1393
+
1394
+
1395
+ _PAULI_MATRICES = {
1396
+ 'I': np.eye(2, dtype=np.complex128),
1397
+ 'X': np.array([[0, 1], [1, 0]], dtype=np.complex128),
1398
+ 'Y': np.array([[0, -1j], [1j, 0]], dtype=np.complex128),
1399
+ 'Z': np.array([[1, 0], [0, -1]], dtype=np.complex128),
1400
+ }
1401
+
1402
+
1403
+ def _mps_transfer_sweep(
1404
+ mps: "MPSSimulator", assignment: dict, need_norm: bool = True
1405
+ ) -> Tuple[complex, Optional[float]]:
1406
+ """Left-to-right transfer-matrix sweep over the Gamma/Lambda tensors
1407
+ computing <psi|P|psi> (unnormalized) for the per-site operator
1408
+ assignment (site -> 'I'|'X'|'Y'|'Z', missing sites default to 'I') --
1409
+ never materializes a (2**n,) statevector, cost O(n * chi^3). When
1410
+ need_norm is True, <psi|psi> is accumulated in the same Python loop
1411
+ (reusing the per-site A = Lambda*Gamma tensor and its conjugate) and
1412
+ returned alongside the raw value, instead of a second, separate sweep
1413
+ -- when False (used by mps_pauli_sum_expectation, which only needs
1414
+ the norm once for the whole sum, not once per term) that accumulation
1415
+ is skipped entirely.
1416
+
1417
+ Every site (not only ones in the assignment) goes through the
1418
+ identical per-site contraction, using the identity matrix wherever the
1419
+ assignment doesn't specify a Pauli -- this is already O(chi^3) per
1420
+ site regardless of whether the local operator is I or a real Pauli, so
1421
+ there is no separate "skip identity sites" fast path to get right or
1422
+ wrong: the same sweep is correct for any placement of the operator's
1423
+ support, including support that spans much of the chain, without
1424
+ relying on which side of any notional orthogonality center a given
1425
+ site falls on.
1426
+
1427
+ The norm is needed at all because SVD truncation (apply_gate_2q) can
1428
+ leave the Gamma/Lambda tensors at less than exact unit norm --
1429
+ contract_to_statevector renormalizes explicitly for the same reason.
1430
+ If the MPS ever carried exact unit norm this division is a no-op, but
1431
+ every reader of the state normalizes explicitly rather than assuming
1432
+ that invariant holds.
1433
+ """
1434
+ dtype = mps.gammas[0].dtype
1435
+ env_op = jnp.ones((1, 1), dtype=dtype)
1436
+ env_norm = jnp.ones((1, 1), dtype=dtype) if need_norm else None
1437
+ for i in range(mps.n):
1438
+ op = _PAULI_MATRICES[assignment.get(i, 'I')].astype(dtype)
1439
+ a = jnp.einsum('l,lpr->lpr', mps.lambdas[i].astype(dtype), mps.gammas[i])
1440
+ a_conj = jnp.conj(a)
1441
+ a_op = jnp.einsum('pq,lqr->lpr', op, a)
1442
+ env_op = jnp.einsum('lk,lpr->kpr', env_op, a_conj)
1443
+ env_op = jnp.einsum('kpr,kps->rs', env_op, a_op)
1444
+ if need_norm:
1445
+ env_norm = jnp.einsum('lk,lpr->kpr', env_norm, a_conj)
1446
+ env_norm = jnp.einsum('kpr,kps->rs', env_norm, a)
1447
+ raw = complex(env_op[0, 0])
1448
+ norm_sq = complex(env_norm[0, 0]).real if need_norm else None
1449
+ return raw, norm_sq
1450
+
1451
+
1452
+ def mps_pauli_expectation(mps: "MPSSimulator", pauli_terms) -> complex:
1453
+ """<psi|P|psi> / <psi|psi> for a single Pauli string P, contracted
1454
+ directly against the MPS (Gamma/Lambda tensors).
1455
+
1456
+ pauli_terms accepts the same three forms as
1457
+ physics.observables.pauli_expectation (a string, e.g. 'XIZ'; a dict
1458
+ {qubit: 'X'|'Y'|'Z'}; or an iterable of (qubit, pauli) pairs) -- reuses
1459
+ that module's own `_normalize_terms` so both functions agree on
1460
+ parsing by construction, not by parallel reimplementation. See
1461
+ _mps_transfer_sweep for why the division is needed.
1462
+ """
1463
+ assignment = _normalize_terms(pauli_terms, mps.n)
1464
+ raw, norm_sq = _mps_transfer_sweep(mps, assignment, need_norm=True)
1465
+ return raw / norm_sq
1466
+
1467
+
1468
+ def mps_pauli_sum_expectation(mps: "MPSSimulator", terms) -> complex:
1469
+ """sum_i coeff_i * <psi|P_i|psi> / <psi|psi> -- same terms format as
1470
+ physics.observables.pauli_sum_expectation: an iterable of
1471
+ (coeff, pauli_terms) pairs. <psi|psi> does not depend on which Pauli
1472
+ string is being measured, so it is computed once via its own sweep
1473
+ and applied to the whole sum, instead of once per term."""
1474
+ terms = list(terms)
1475
+ if not terms:
1476
+ return 0j
1477
+ _, norm_sq = _mps_transfer_sweep(mps, {}, need_norm=True)
1478
+ raw_sum = sum(
1479
+ coeff * _mps_transfer_sweep(mps, _normalize_terms(pauli_terms, mps.n), need_norm=False)[0]
1480
+ for coeff, pauli_terms in terms
1481
+ )
1482
+ return raw_sum / norm_sq
1483
+
1484
+
1485
+ @dataclasses.dataclass
1486
+ class BondConvergenceResult:
1487
+ bonds: List[int]
1488
+ chi_used: List[int]
1489
+ avg_jsd: List[float]
1490
+ budget_violations: List[int]
1491
+ values: List[List[complex]]
1492
+ diffs: List[List[float]]
1493
+ verdicts: List[str]
1494
+
1495
+
1496
+ def bond_convergence(
1497
+ ops: List, n_qubits: int, observables: list, bonds: List[int],
1498
+ tol: float = 1e-3, **mps_kwargs,
1499
+ ) -> BondConvergenceResult:
1500
+ """Runs the same circuit at every value in `bonds` (increasing) and
1501
+ checks whether the reported observables have actually converged with
1502
+ respect to bond dimension, instead of trusting a single run's own
1503
+ internal diagnostics.
1504
+
1505
+ Requires len(bonds) >= 3. Two bonds give exactly one discrepancy,
1506
+ which is a single number with no way to tell whether it is still
1507
+ shrinking toward `tol` or has already stalled -- measured on a
1508
+ 40-qubit, 4-layer brickwall circuit, chi=4->8->32 gave |<Z0>|
1509
+ discrepancies of ~4.7e-2 then ~1.2e-2 (chi_used never hit its own
1510
+ cap, so this is a real not_converged, not an artifact of running out
1511
+ of bond dimension): a two-bond check (chi=4 vs 8) would see only the
1512
+ first number and have no basis to call it anything, while three
1513
+ bonds show a trend that is decreasing but still two orders of
1514
+ magnitude above any reasonable `tol`.
1515
+
1516
+ A verdict of "converged" additionally requires the successive
1517
+ discrepancies to be monotonically non-increasing, not just that the
1518
+ last one is below `tol` -- a single small discrepancy proves nothing
1519
+ about the trend on its own, which is the same failure mode as the
1520
+ two-bond case above, one level up. (Ties count as non-increasing: an
1521
+ exactly-converged observable, e.g. a GHZ chain whose bond dimension
1522
+ never needs to grow, produces identical values -- and therefore
1523
+ zero discrepancies -- at every bond, which must count as converged.)
1524
+
1525
+ avg_jsd and budget_violations (from the underlying MPSSimulator runs)
1526
+ are reported per bond for context only, never used to decide the
1527
+ verdict -- a low average JSD is computed per truncation step and says
1528
+ nothing about whether the specific observable being tracked has
1529
+ settled down as `max_bond` grows.
1530
+
1531
+ If max_bond_used() at the highest bond still equals that bond's cap,
1532
+ the truncation never had headroom below max_bond at any cut, so no
1533
+ tolerance can be certified from this data: every observable's verdict
1534
+ becomes "undecidable" regardless of its own discrepancies.
1535
+ """
1536
+ if len(bonds) < 3:
1537
+ raise ValueError(f"bond_convergence needs at least 3 bonds to detect a trend, got {len(bonds)}")
1538
+
1539
+ chi_used, avg_jsd, budget_violations = [], [], []
1540
+ values = [[] for _ in observables]
1541
+ for bond in bonds:
1542
+ mps = MPSSimulator(n_qubits=n_qubits, max_bond=bond, **mps_kwargs)
1543
+ mps.run_circuit_jit(ops)
1544
+ chi_used.append(mps.max_bond_used())
1545
+ avg_jsd.append(mps.avg_jsd())
1546
+ budget_violations.append(mps.budget_violations)
1547
+ for obs_idx, obs in enumerate(observables):
1548
+ values[obs_idx].append(mps_pauli_expectation(mps, obs))
1549
+
1550
+ diffs = [
1551
+ [abs(vals[i + 1] - vals[i]) for i in range(len(vals) - 1)]
1552
+ for vals in values
1553
+ ]
1554
+
1555
+ undecidable = chi_used[-1] >= bonds[-1]
1556
+ verdicts = []
1557
+ for d in diffs:
1558
+ if undecidable:
1559
+ verdicts.append("undecidable")
1560
+ elif all(d[i + 1] <= d[i] for i in range(len(d) - 1)) and d[-1] < tol:
1561
+ verdicts.append("converged")
1562
+ else:
1563
+ verdicts.append("not_converged")
1564
+
1565
+ return BondConvergenceResult(
1566
+ bonds=list(bonds), chi_used=chi_used, avg_jsd=avg_jsd,
1567
+ budget_violations=budget_violations, values=values, diffs=diffs,
1568
+ verdicts=verdicts,
1569
+ )