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.
- dashboard_core/__init__.py +115 -0
- dashboard_core/_gate_tables.py +30 -0
- dashboard_core/band_structure.py +71 -0
- dashboard_core/circuit_builder_component.py +232 -0
- dashboard_core/circuit_diagram.py +216 -0
- dashboard_core/crypto_protocols.py +77 -0
- dashboard_core/engine.py +326 -0
- dashboard_core/graphical_builder.py +114 -0
- dashboard_core/hamiltonians.py +593 -0
- dashboard_core/mass_decomposition_tool.py +47 -0
- dashboard_core/mitigation.py +343 -0
- dashboard_core/native_hf_diagnostics.py +62 -0
- dashboard_core/noise_tools.py +125 -0
- dashboard_core/qasm_library.py +233 -0
- dashboard_core/qmmm.py +16 -0
- dashboard_core/rag_tool.py +45 -0
- dashboard_core/state_visuals.py +288 -0
- dashboard_core/system_limits.py +60 -0
- dashboard_core/vector_healing.py +102 -0
- dashboard_core/visuals.py +158 -0
- dashboard_core/vqe.py +533 -0
- dashboard_core/wormhole.py +580 -0
- dense_evolution/__init__.py +114 -0
- dense_evolution/autodiff.py +10 -0
- dense_evolution/backends/__init__.py +5 -0
- dense_evolution/backends/chunk/__init__.py +37 -0
- dense_evolution/backends/chunk/_engine_imports.py +57 -0
- dense_evolution/backends/chunk/circuit_chunker.py +55 -0
- dense_evolution/backends/chunk/core.py +432 -0
- dense_evolution/backends/chunk/disk_overflow.py +232 -0
- dense_evolution/backends/chunk/geometry.py +95 -0
- dense_evolution/backends/chunk/guard.py +190 -0
- dense_evolution/backends/chunk/kernels.py +531 -0
- dense_evolution/backends/mps.py +1569 -0
- dense_evolution/backends/statevector.py +616 -0
- dense_evolution/chunk.py +25 -0
- dense_evolution/circuits/__init__.py +20 -0
- dense_evolution/circuits/compiler.py +488 -0
- dense_evolution/circuits/diagram.py +94 -0
- dense_evolution/circuits/gates.py +91 -0
- dense_evolution/circuits/parser.py +632 -0
- dense_evolution/circuits/qft.py +66 -0
- dense_evolution/circuits/random_circuit.py +85 -0
- dense_evolution/circuits/registry.py +74 -0
- dense_evolution/circuits/topology.py +79 -0
- dense_evolution/circuits/trotter.py +265 -0
- dense_evolution/circuits/uccsd.py +275 -0
- dense_evolution/cli.py +199 -0
- dense_evolution/compiler.py +9 -0
- dense_evolution/config.py +49 -0
- dense_evolution/drawing.py +10 -0
- dense_evolution/entropy.py +9 -0
- dense_evolution/fermions.py +9 -0
- dense_evolution/gates.py +9 -0
- dense_evolution/harrison_tb.py +16 -0
- dense_evolution/healing.py +18 -0
- dense_evolution/interop/__init__.py +18 -0
- dense_evolution/interop/qiskit_pennylane.py +406 -0
- dense_evolution/measurement.py +10 -0
- dense_evolution/mitigation/__init__.py +54 -0
- dense_evolution/mitigation/healing.py +215 -0
- dense_evolution/mitigation/kl_divergence.py +93 -0
- dense_evolution/mitigation/magic_entropy.py +163 -0
- dense_evolution/mitigation/magic_entropy_shadows.py +262 -0
- dense_evolution/mitigation/renyi.py +168 -0
- dense_evolution/mitigation/stabilizer_renyi_entropy.py +103 -0
- dense_evolution/mitigation/zne.py +990 -0
- dense_evolution/mps.py +9 -0
- dense_evolution/native_hf/__init__.py +26 -0
- dense_evolution/native_hf/_libcint/LICENSE-libcint +10 -0
- dense_evolution/native_hf/_libcint/libdecint.dll +0 -0
- dense_evolution/native_hf/assembly.py +304 -0
- dense_evolution/native_hf/basis.py +117 -0
- dense_evolution/native_hf/boys.py +35 -0
- dense_evolution/native_hf/bridge.py +112 -0
- dense_evolution/native_hf/cartesian.py +64 -0
- dense_evolution/native_hf/coulomb.py +196 -0
- dense_evolution/native_hf/differentiable.py +53 -0
- dense_evolution/native_hf/gaussians.py +79 -0
- dense_evolution/native_hf/kinetic.py +52 -0
- dense_evolution/native_hf/libcint_bridge.py +167 -0
- dense_evolution/native_hf/overlap.py +91 -0
- dense_evolution/native_hf/scf.py +404 -0
- dense_evolution/noise/__init__.py +79 -0
- dense_evolution/noise/coherent_attack.py +264 -0
- dense_evolution/noise/cosmic_ray.py +61 -0
- dense_evolution/noise/density_matrix_channels.py +78 -0
- dense_evolution/noise/differentiable.py +66 -0
- dense_evolution/noise/kraus/__init__.py +6 -0
- dense_evolution/noise/kraus/amplitude_damping.py +47 -0
- dense_evolution/noise/kraus/bitflip.py +22 -0
- dense_evolution/noise/kraus/combined.py +16 -0
- dense_evolution/noise/kraus/depolarizing.py +47 -0
- dense_evolution/noise/kraus/ideal.py +10 -0
- dense_evolution/noise/kraus/phaseflip.py +21 -0
- dense_evolution/noise/kraus_channels.py +285 -0
- dense_evolution/noise/oscillating.py +32 -0
- dense_evolution/noise/pink.py +80 -0
- dense_evolution/observables.py +11 -0
- dense_evolution/parser.py +9 -0
- dense_evolution/physics/__init__.py +27 -0
- dense_evolution/physics/entropy.py +161 -0
- dense_evolution/physics/fermions.py +322 -0
- dense_evolution/physics/observables.py +523 -0
- dense_evolution/physics/qec.py +1113 -0
- dense_evolution/physics/spectral.py +143 -0
- dense_evolution/physics/states.py +43 -0
- dense_evolution/protocols/__init__.py +27 -0
- dense_evolution/protocols/bb84.py +133 -0
- dense_evolution/protocols/di_qkd_ghz.py +199 -0
- dense_evolution/protocols/dicka_protocol2.py +124 -0
- dense_evolution/qec.py +20 -0
- dense_evolution/qft.py +9 -0
- dense_evolution/qmmm/__init__.py +13 -0
- dense_evolution/qmmm/ase_bridge.py +97 -0
- dense_evolution/qmmm/forces.py +388 -0
- dense_evolution/qmmm/propagation.py +80 -0
- dense_evolution/qmmm/region.py +137 -0
- dense_evolution/random_circuit.py +15 -0
- dense_evolution/registry.py +9 -0
- dense_evolution/simulator.py +10 -0
- dense_evolution/solvers/__init__.py +19 -0
- dense_evolution/solvers/autodiff.py +169 -0
- dense_evolution/solvers/harrison_tb.py +189 -0
- dense_evolution/solvers/vhd_tb.py +187 -0
- dense_evolution/states.py +9 -0
- dense_evolution/topology.py +9 -0
- dense_evolution/trotter.py +9 -0
- dense_evolution/utils/__init__.py +13 -0
- dense_evolution/utils/drawing.py +101 -0
- dense_evolution/utils/mass_decomposition.py +246 -0
- dense_evolution/utils/measurement.py +94 -0
- dense_evolution/vhd_tb.py +16 -0
- dense_evolution-8.3.0.dist-info/METADATA +366 -0
- dense_evolution-8.3.0.dist-info/RECORD +165 -0
- dense_evolution-8.3.0.dist-info/WHEEL +5 -0
- dense_evolution-8.3.0.dist-info/entry_points.txt +2 -0
- dense_evolution-8.3.0.dist-info/licenses/license.md +58 -0
- dense_evolution-8.3.0.dist-info/top_level.txt +5 -0
- ia_utils/__init__.py +0 -0
- ia_utils/adversarial_vector_attack.py +196 -0
- ia_utils/rag.py +288 -0
- ia_utils/vector_healing.py +399 -0
- local_site/__init__.py +0 -0
- local_site/app/__init__.py +0 -0
- local_site/app/server.py +1009 -0
- mcp_server/__init__.py +0 -0
- mcp_server/client.py +324 -0
- mcp_server/config.py +32 -0
- mcp_server/models.py +347 -0
- mcp_server/molecules.py +71 -0
- mcp_server/server.py +119 -0
- mcp_server/tools/__init__.py +0 -0
- mcp_server/tools/chemistry_tools.py +225 -0
- mcp_server/tools/circuit_tools.py +83 -0
- mcp_server/tools/crypto_tools.py +66 -0
- mcp_server/tools/mitigation_tools.py +81 -0
- mcp_server/tools/noise_tools.py +60 -0
- mcp_server/tools/retrieval_tools.py +44 -0
- mcp_server/tools/system_tools.py +149 -0
- mcp_server/tools/wormhole_tools.py +142 -0
- mcp_server/utils/__init__.py +0 -0
- mcp_server/utils/cache.py +55 -0
- mcp_server/utils/images.py +67 -0
- 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
|
+
)
|