dense-evolution 8.3.0__py3-none-win_amd64.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (165) hide show
  1. dashboard_core/__init__.py +115 -0
  2. dashboard_core/_gate_tables.py +30 -0
  3. dashboard_core/band_structure.py +71 -0
  4. dashboard_core/circuit_builder_component.py +232 -0
  5. dashboard_core/circuit_diagram.py +216 -0
  6. dashboard_core/crypto_protocols.py +77 -0
  7. dashboard_core/engine.py +326 -0
  8. dashboard_core/graphical_builder.py +114 -0
  9. dashboard_core/hamiltonians.py +593 -0
  10. dashboard_core/mass_decomposition_tool.py +47 -0
  11. dashboard_core/mitigation.py +343 -0
  12. dashboard_core/native_hf_diagnostics.py +62 -0
  13. dashboard_core/noise_tools.py +125 -0
  14. dashboard_core/qasm_library.py +233 -0
  15. dashboard_core/qmmm.py +16 -0
  16. dashboard_core/rag_tool.py +45 -0
  17. dashboard_core/state_visuals.py +288 -0
  18. dashboard_core/system_limits.py +60 -0
  19. dashboard_core/vector_healing.py +102 -0
  20. dashboard_core/visuals.py +158 -0
  21. dashboard_core/vqe.py +533 -0
  22. dashboard_core/wormhole.py +580 -0
  23. dense_evolution/__init__.py +114 -0
  24. dense_evolution/autodiff.py +10 -0
  25. dense_evolution/backends/__init__.py +5 -0
  26. dense_evolution/backends/chunk/__init__.py +37 -0
  27. dense_evolution/backends/chunk/_engine_imports.py +57 -0
  28. dense_evolution/backends/chunk/circuit_chunker.py +55 -0
  29. dense_evolution/backends/chunk/core.py +432 -0
  30. dense_evolution/backends/chunk/disk_overflow.py +232 -0
  31. dense_evolution/backends/chunk/geometry.py +95 -0
  32. dense_evolution/backends/chunk/guard.py +190 -0
  33. dense_evolution/backends/chunk/kernels.py +531 -0
  34. dense_evolution/backends/mps.py +1569 -0
  35. dense_evolution/backends/statevector.py +616 -0
  36. dense_evolution/chunk.py +25 -0
  37. dense_evolution/circuits/__init__.py +20 -0
  38. dense_evolution/circuits/compiler.py +488 -0
  39. dense_evolution/circuits/diagram.py +94 -0
  40. dense_evolution/circuits/gates.py +91 -0
  41. dense_evolution/circuits/parser.py +632 -0
  42. dense_evolution/circuits/qft.py +66 -0
  43. dense_evolution/circuits/random_circuit.py +85 -0
  44. dense_evolution/circuits/registry.py +74 -0
  45. dense_evolution/circuits/topology.py +79 -0
  46. dense_evolution/circuits/trotter.py +265 -0
  47. dense_evolution/circuits/uccsd.py +275 -0
  48. dense_evolution/cli.py +199 -0
  49. dense_evolution/compiler.py +9 -0
  50. dense_evolution/config.py +49 -0
  51. dense_evolution/drawing.py +10 -0
  52. dense_evolution/entropy.py +9 -0
  53. dense_evolution/fermions.py +9 -0
  54. dense_evolution/gates.py +9 -0
  55. dense_evolution/harrison_tb.py +16 -0
  56. dense_evolution/healing.py +18 -0
  57. dense_evolution/interop/__init__.py +18 -0
  58. dense_evolution/interop/qiskit_pennylane.py +406 -0
  59. dense_evolution/measurement.py +10 -0
  60. dense_evolution/mitigation/__init__.py +54 -0
  61. dense_evolution/mitigation/healing.py +215 -0
  62. dense_evolution/mitigation/kl_divergence.py +93 -0
  63. dense_evolution/mitigation/magic_entropy.py +163 -0
  64. dense_evolution/mitigation/magic_entropy_shadows.py +262 -0
  65. dense_evolution/mitigation/renyi.py +168 -0
  66. dense_evolution/mitigation/stabilizer_renyi_entropy.py +103 -0
  67. dense_evolution/mitigation/zne.py +990 -0
  68. dense_evolution/mps.py +9 -0
  69. dense_evolution/native_hf/__init__.py +26 -0
  70. dense_evolution/native_hf/_libcint/LICENSE-libcint +10 -0
  71. dense_evolution/native_hf/_libcint/libdecint.dll +0 -0
  72. dense_evolution/native_hf/assembly.py +304 -0
  73. dense_evolution/native_hf/basis.py +117 -0
  74. dense_evolution/native_hf/boys.py +35 -0
  75. dense_evolution/native_hf/bridge.py +112 -0
  76. dense_evolution/native_hf/cartesian.py +64 -0
  77. dense_evolution/native_hf/coulomb.py +196 -0
  78. dense_evolution/native_hf/differentiable.py +53 -0
  79. dense_evolution/native_hf/gaussians.py +79 -0
  80. dense_evolution/native_hf/kinetic.py +52 -0
  81. dense_evolution/native_hf/libcint_bridge.py +167 -0
  82. dense_evolution/native_hf/overlap.py +91 -0
  83. dense_evolution/native_hf/scf.py +404 -0
  84. dense_evolution/noise/__init__.py +79 -0
  85. dense_evolution/noise/coherent_attack.py +264 -0
  86. dense_evolution/noise/cosmic_ray.py +61 -0
  87. dense_evolution/noise/density_matrix_channels.py +78 -0
  88. dense_evolution/noise/differentiable.py +66 -0
  89. dense_evolution/noise/kraus/__init__.py +6 -0
  90. dense_evolution/noise/kraus/amplitude_damping.py +47 -0
  91. dense_evolution/noise/kraus/bitflip.py +22 -0
  92. dense_evolution/noise/kraus/combined.py +16 -0
  93. dense_evolution/noise/kraus/depolarizing.py +47 -0
  94. dense_evolution/noise/kraus/ideal.py +10 -0
  95. dense_evolution/noise/kraus/phaseflip.py +21 -0
  96. dense_evolution/noise/kraus_channels.py +285 -0
  97. dense_evolution/noise/oscillating.py +32 -0
  98. dense_evolution/noise/pink.py +80 -0
  99. dense_evolution/observables.py +11 -0
  100. dense_evolution/parser.py +9 -0
  101. dense_evolution/physics/__init__.py +27 -0
  102. dense_evolution/physics/entropy.py +161 -0
  103. dense_evolution/physics/fermions.py +322 -0
  104. dense_evolution/physics/observables.py +523 -0
  105. dense_evolution/physics/qec.py +1113 -0
  106. dense_evolution/physics/spectral.py +143 -0
  107. dense_evolution/physics/states.py +43 -0
  108. dense_evolution/protocols/__init__.py +27 -0
  109. dense_evolution/protocols/bb84.py +133 -0
  110. dense_evolution/protocols/di_qkd_ghz.py +199 -0
  111. dense_evolution/protocols/dicka_protocol2.py +124 -0
  112. dense_evolution/qec.py +20 -0
  113. dense_evolution/qft.py +9 -0
  114. dense_evolution/qmmm/__init__.py +13 -0
  115. dense_evolution/qmmm/ase_bridge.py +97 -0
  116. dense_evolution/qmmm/forces.py +388 -0
  117. dense_evolution/qmmm/propagation.py +80 -0
  118. dense_evolution/qmmm/region.py +137 -0
  119. dense_evolution/random_circuit.py +15 -0
  120. dense_evolution/registry.py +9 -0
  121. dense_evolution/simulator.py +10 -0
  122. dense_evolution/solvers/__init__.py +19 -0
  123. dense_evolution/solvers/autodiff.py +169 -0
  124. dense_evolution/solvers/harrison_tb.py +189 -0
  125. dense_evolution/solvers/vhd_tb.py +187 -0
  126. dense_evolution/states.py +9 -0
  127. dense_evolution/topology.py +9 -0
  128. dense_evolution/trotter.py +9 -0
  129. dense_evolution/utils/__init__.py +13 -0
  130. dense_evolution/utils/drawing.py +101 -0
  131. dense_evolution/utils/mass_decomposition.py +246 -0
  132. dense_evolution/utils/measurement.py +94 -0
  133. dense_evolution/vhd_tb.py +16 -0
  134. dense_evolution-8.3.0.dist-info/METADATA +366 -0
  135. dense_evolution-8.3.0.dist-info/RECORD +165 -0
  136. dense_evolution-8.3.0.dist-info/WHEEL +5 -0
  137. dense_evolution-8.3.0.dist-info/entry_points.txt +2 -0
  138. dense_evolution-8.3.0.dist-info/licenses/license.md +58 -0
  139. dense_evolution-8.3.0.dist-info/top_level.txt +5 -0
  140. ia_utils/__init__.py +0 -0
  141. ia_utils/adversarial_vector_attack.py +196 -0
  142. ia_utils/rag.py +288 -0
  143. ia_utils/vector_healing.py +399 -0
  144. local_site/__init__.py +0 -0
  145. local_site/app/__init__.py +0 -0
  146. local_site/app/server.py +1009 -0
  147. mcp_server/__init__.py +0 -0
  148. mcp_server/client.py +324 -0
  149. mcp_server/config.py +32 -0
  150. mcp_server/models.py +347 -0
  151. mcp_server/molecules.py +71 -0
  152. mcp_server/server.py +119 -0
  153. mcp_server/tools/__init__.py +0 -0
  154. mcp_server/tools/chemistry_tools.py +225 -0
  155. mcp_server/tools/circuit_tools.py +83 -0
  156. mcp_server/tools/crypto_tools.py +66 -0
  157. mcp_server/tools/mitigation_tools.py +81 -0
  158. mcp_server/tools/noise_tools.py +60 -0
  159. mcp_server/tools/retrieval_tools.py +44 -0
  160. mcp_server/tools/system_tools.py +149 -0
  161. mcp_server/tools/wormhole_tools.py +142 -0
  162. mcp_server/utils/__init__.py +0 -0
  163. mcp_server/utils/cache.py +55 -0
  164. mcp_server/utils/images.py +67 -0
  165. mcp_server/utils/truncation.py +38 -0
@@ -0,0 +1,990 @@
1
+ """
2
+ Zero-Noise Extrapolation (ZNE)
3
+ -------------------------------
4
+ Standard error-mitigation entry points, named the way the field already
5
+ names them (Richardson extrapolation, noise factors, zero-noise
6
+ extrapolation -- same vocabulary as e.g. Mitiq's `zne` API), so callers
7
+ and tooling can find "ZNE" without first learning Dense-Evolution's
8
+ internal healing vocabulary.
9
+
10
+ This module composes `dense_evolution.healing`'s existing primitives
11
+ (`calculate_delta_preemp`, ...) -- it does not rename or replace them.
12
+ """
13
+ import functools
14
+ import warnings
15
+
16
+ import numpy as np
17
+ import jax
18
+ import jax.numpy as jnp
19
+ from scipy.optimize import minimize
20
+
21
+ from .healing import calculate_delta_preemp
22
+
23
+ __all__ = ["richardson_extrapolate", "richardson_amplification_factor", "zero_noise_extrapolation",
24
+ "polynomial_extrapolate", "bounded_exponential_extrapolate",
25
+ "project_to_physical", "uhlmann_fidelity", "zne_density_matrix",
26
+ "jsd_predictive_zne_density_matrix", "coherence_predictive_zne_density_matrix",
27
+ "classically_augmented_zne_phaseflip", "global_depolarizing_channel",
28
+ "amplitude_damping_channel", "cosmic_ray_burst_profile",
29
+ "richardson_extrapolate_jit", "zero_noise_extrapolation_jit",
30
+ "polynomial_extrapolate_jit", "uhlmann_fidelity_jit", "zne_density_matrix_jit"]
31
+
32
+ _KAPPA_WARNING_THRESHOLD = 50.0
33
+ _MULTI_START_SEED = 0
34
+ _N_MULTI_STARTS = 8
35
+
36
+
37
+ def richardson_extrapolate(expectation_values, noise_factors) -> jnp.ndarray:
38
+ """Polynomial (Lagrange) Richardson extrapolation to zero noise.
39
+
40
+ `expectation_values[i]` is the value measured/simulated at noise scale
41
+ `noise_factors[i]` (e.g. 1x, 2x, 3x folded/scaled noise) -- a scalar,
42
+ or itself an array (e.g. a full probability distribution sampled at
43
+ that noise scale; extrapolated elementwise). Returns the extrapolated
44
+ zero-noise estimate, same shape as one `expectation_values[i]`. Works
45
+ for any number of points and any (not necessarily equally spaced)
46
+ noise factors; for the common 3-point case at noise_factors=(1,2,3)
47
+ this reduces exactly to the textbook coefficients (3, -3, 1).
48
+
49
+ `expectation_values` may be complex (e.g. density matrix entries,
50
+ which are complex off-diagonal in general) -- dtype is picked from
51
+ the input itself (complex128 if complex, float64 otherwise, matching
52
+ this function's previous always-float64 behavior for real input
53
+ exactly). Found via a real test case: forcing float64 unconditionally
54
+ silently discarded the imaginary part of complex input with no
55
+ visible error, only a low-signal ComplexWarning easy to miss --
56
+ confirmed directly (`richardson_extrapolate([1+2j, 3+4j], ...)` used
57
+ to return a purely real result, dropping real information).
58
+
59
+ Passing more noise-scale points makes exact interpolation MORE, not
60
+ less, sensitive to shot noise: at noise_factors equally spaced in
61
+ [1, 3], the noise-amplification factor kappa (`richardson_amplification_factor`,
62
+ the sum of absolute Lagrange coefficients) is 7 at 3 points, 129 at 5,
63
+ 2815 at 7 -- verified directly, exact integers, not estimates. Since
64
+ the coefficients are fixed constants (independent of the measured
65
+ values), i.i.d. shot noise of per-point standard deviation sigma
66
+ propagates to an extrapolated-result standard deviation of
67
+ `sigma * sqrt(sum(coeff_i**2))`: at sigma=0.01 this is 0.044 (3
68
+ points), 0.67 (5), 13.2 (7) -- verified with this exact formula, not
69
+ simulated, and matching a real prior audit's Monte Carlo figures
70
+ (0.044 / 0.67 / 13.5) to within sampling noise. Against this,
71
+ prog.txt's own audit reports the systematic (noise-free) bias
72
+ improving by only about 0.01 over the same 3-to-7-point range on its
73
+ real experimental setup -- a bad trade past a handful of points. A
74
+ `UserWarning` fires when kappa exceeds `_KAPPA_WARNING_THRESHOLD`
75
+ (50); see `richardson_amplification_factor` to check kappa before
76
+ extrapolating.
77
+
78
+ Examples
79
+ --------
80
+ >>> from dense_evolution.mitigation import richardson_extrapolate
81
+ >>> round(float(richardson_extrapolate([0.90, 0.80, 0.65], [1.0, 2.0, 3.0])), 4)
82
+ 0.95
83
+ """
84
+ lambdas = jnp.asarray(noise_factors, dtype=jnp.float64)
85
+ kappa = richardson_amplification_factor(lambdas)
86
+ if kappa > _KAPPA_WARNING_THRESHOLD:
87
+ warnings.warn(
88
+ f"richardson_extrapolate: noise amplification factor kappa={kappa:.1f} exceeds "
89
+ f"{_KAPPA_WARNING_THRESHOLD:.0f} -- shot noise in expectation_values will be "
90
+ f"amplified by roughly this factor in the extrapolated result; consider passing "
91
+ f"fewer noise_factors points.",
92
+ UserWarning, stacklevel=2,
93
+ )
94
+ values_dtype = jnp.complex128 if np.iscomplexobj(np.asarray(expectation_values)) else jnp.float64
95
+ values = jnp.asarray(expectation_values, dtype=values_dtype)
96
+ return _richardson_extrapolate_core(values, lambdas)
97
+
98
+
99
+ def _lagrange_coeffs_at_zero(lambdas: jnp.ndarray) -> jnp.ndarray:
100
+ """Lagrange basis polynomials for nodes `lambdas`, each evaluated at
101
+ x=0 -- shared by `_richardson_extrapolate_core` (which combines them
102
+ with the measured values) and `richardson_amplification_factor`
103
+ (which only needs their magnitudes)."""
104
+ n = lambdas.shape[0]
105
+
106
+ def lagrange_coeff(i):
107
+ others = jnp.concatenate([lambdas[:i], lambdas[i + 1:]])
108
+ return jnp.prod((0.0 - others) / (lambdas[i] - others))
109
+
110
+ return jnp.stack([lagrange_coeff(i) for i in range(n)])
111
+
112
+
113
+ def richardson_amplification_factor(noise_factors) -> float:
114
+ """Noise-amplification factor kappa = sum(|Lagrange coefficients|) for
115
+ Richardson extrapolation at the given `noise_factors`: how much a
116
+ unit of i.i.d. shot noise spread evenly across the measured points
117
+ gets amplified in the extrapolated zero-noise estimate (each
118
+ coefficient can be much larger than 1 and they alternate in sign, so
119
+ they do not cancel in the worst case the way their SUM, which is
120
+ always exactly 1, might suggest).
121
+
122
+ Measured at noise_factors equally spaced in [1, 3]: kappa = 7 at 3
123
+ points, 129 at 5 points, 2815 at 7 points -- see
124
+ `richardson_extrapolate`'s own docstring for the resulting bias/
125
+ variance trade-off in these same three cases.
126
+
127
+ Examples
128
+ --------
129
+ >>> from dense_evolution.mitigation import richardson_amplification_factor
130
+ >>> round(richardson_amplification_factor([1.0, 2.0, 3.0]), 4)
131
+ 7.0
132
+ """
133
+ lambdas = jnp.asarray(noise_factors, dtype=jnp.float64)
134
+ coeffs = _lagrange_coeffs_at_zero(lambdas)
135
+ return float(jnp.sum(jnp.abs(coeffs)))
136
+
137
+
138
+ def _richardson_extrapolate_core(values: jnp.ndarray, lambdas: jnp.ndarray) -> jnp.ndarray:
139
+ """`jax.jit`-traceable core of `richardson_extrapolate` -- `values`
140
+ already cast to its final dtype, `lambdas` already float64 (no
141
+ `np.iscomplexobj`/`np.asarray` call on a possibly-traced value). `n`
142
+ comes from `lambdas.shape[0]`, a static value even under tracing
143
+ (array shapes are always known at trace time), so the Python
144
+ `range(n)` unroll below is trace-safe without `n` needing to be
145
+ static-marked by callers.
146
+ """
147
+ n = lambdas.shape[0]
148
+ coeffs = _lagrange_coeffs_at_zero(lambdas)
149
+ # Broadcast coeffs against the LEADING axis of values (the "one row per
150
+ # noise scale" axis), not jnp's default trailing-axis alignment --
151
+ # values may itself be array-valued per scale (values.shape = (n,
152
+ # *extra_dims)), e.g. a whole probability distribution rather than a
153
+ # bare scalar. A no-op reshape when values is 1-D (the scalar case),
154
+ # so existing scalar callers are unaffected.
155
+ coeffs = coeffs.reshape((n,) + (1,) * (values.ndim - 1))
156
+ return jnp.sum(coeffs * values, axis=0)
157
+
158
+
159
+ richardson_extrapolate_jit = jax.jit(_richardson_extrapolate_core)
160
+ """`jax.jit`-compiled entry point for `richardson_extrapolate`. Unlike
161
+ `polynomial_extrapolate`/`zne_density_matrix`'s jitted variants, no
162
+ argument needs to be marked static here -- `n` (the point count) is read
163
+ from `lambdas.shape[0]`, itself always static under tracing, not from a
164
+ Python `degree` parameter.
165
+
166
+ `values` must already be complex128 or float64 (pick the dtype yourself
167
+ before calling -- this skips `richardson_extrapolate`'s `np.iscomplexobj`
168
+ auto-detection, which isn't traceable) and `lambdas` a float64 array.
169
+ Verified to match `richardson_extrapolate` exactly on real and complex
170
+ input; JAX recompiles per distinct input shape/dtype, as usual.
171
+ """
172
+
173
+
174
+ def zero_noise_extrapolation(expectation_values, noise_factors,
175
+ sigma_at_base_noise=None,
176
+ target_sigma_ideal: float = 10.0) -> jnp.ndarray:
177
+ """Zero-Noise Extrapolation -- plain, or healing-adapted when a
178
+ coherence signal is available.
179
+
180
+ Without `sigma_at_base_noise`: standard Richardson ZNE
181
+ (`richardson_extrapolate`).
182
+
183
+ With `sigma_at_base_noise` (the measured/simulated coherence sigma at
184
+ the base, unscaled noise level): the 3 Richardson coefficients are
185
+ perturbed by `dense_evolution.healing.calculate_delta_preemp` -- the
186
+ normalized deviation between the observed sigma and the ideal target
187
+ -- then renormalized to sum to 1. This is Dense-Evolution's
188
+ "predictive healing" ZNE variant: when the observed coherence is off
189
+ the ideal target, the extrapolation is nudged accordingly instead of
190
+ trusting the 3 raw noise-scaled points equally.
191
+
192
+ The healing-adapted path currently only supports exactly 3 noise
193
+ factors (the case it has been derived and tested against); passing
194
+ `sigma_at_base_noise` with any other point count raises
195
+ NotImplementedError rather than silently generalizing an unverified
196
+ formula.
197
+ """
198
+ if sigma_at_base_noise is None:
199
+ return richardson_extrapolate(expectation_values, noise_factors)
200
+
201
+ lambdas = jnp.asarray(noise_factors, dtype=jnp.float64)
202
+ if lambdas.shape[0] != 3:
203
+ raise NotImplementedError(
204
+ "Healing-adapted ZNE is only defined for exactly 3 noise factors; "
205
+ "call richardson_extrapolate(...) directly for the plain N-point case."
206
+ )
207
+ values_dtype = jnp.complex128 if np.iscomplexobj(np.asarray(expectation_values)) else jnp.float64
208
+ values = jnp.asarray(expectation_values, dtype=values_dtype)
209
+ return _zero_noise_extrapolation_healing_core(
210
+ values, jnp.asarray(sigma_at_base_noise, dtype=jnp.float64), target_sigma_ideal)
211
+
212
+
213
+ def _zero_noise_extrapolation_healing_core(values: jnp.ndarray, sigma_at_base_noise: jnp.ndarray,
214
+ target_sigma_ideal: float) -> jnp.ndarray:
215
+ """`jax.jit`-traceable core of `zero_noise_extrapolation`'s
216
+ healing-adapted branch -- `values` already cast to its final dtype
217
+ (exactly 3 noise-scale rows, `values[0]`, `values[1]`, `values[2]`),
218
+ `sigma_at_base_noise` already a jnp float64 scalar. `calculate_delta_preemp`
219
+ is itself already `@jax.jit`-decorated (`dense_evolution.healing`), so
220
+ calling it here composes cleanly under an outer jit.
221
+ """
222
+ e_l1, e_l2, e_l3 = values[0], values[1], values[2]
223
+ delta_p = calculate_delta_preemp(sigma_at_base_noise, target_sigma_ideal)
224
+ c1 = 3.0 - 0.01 * delta_p
225
+ c2 = -3.0 + 0.02 * delta_p
226
+ c3 = 1.0 - 0.01 * delta_p
227
+ return (c1 * e_l1 + c2 * e_l2 + c3 * e_l3) / (c1 + c2 + c3)
228
+
229
+
230
+ zero_noise_extrapolation_jit = jax.jit(_zero_noise_extrapolation_healing_core)
231
+ """`jax.jit`-compiled entry point for `zero_noise_extrapolation`'s
232
+ healing-adapted branch (the `sigma_at_base_noise is not None` case --
233
+ the plain-Richardson case already has `richardson_extrapolate_jit`, use
234
+ that directly instead). `values` must already be complex128 or float64
235
+ with exactly 3 rows, `sigma_at_base_noise` a float64 scalar; the
236
+ `lambdas.shape[0] != 3` validation and dtype auto-detection that
237
+ `zero_noise_extrapolation` does are both skipped here (not traceable) --
238
+ callers are responsible for passing exactly-3-row input themselves.
239
+ `target_sigma_ideal` is a plain Python float, fine to leave non-static
240
+ since it only ever multiplies/subtracts, no Python branching on its value.
241
+ """
242
+
243
+
244
+ def polynomial_extrapolate(expectation_values, noise_factors, degree: int = 2) -> jnp.ndarray:
245
+ """Least-squares polynomial extrapolation to zero noise.
246
+
247
+ Generalizes `richardson_extrapolate`: fits a degree-`degree` polynomial
248
+ to `(noise_factors, expectation_values)` by ordinary least squares and
249
+ evaluates it at zero. With exactly `degree + 1` points the fit is the
250
+ unique interpolating polynomial, mathematically identical to
251
+ `richardson_extrapolate` at that point count (verified directly, both
252
+ for real and complex input). With MORE than `degree + 1` points it
253
+ becomes an overdetermined fit -- the extra points average down
254
+ statistical noise instead of forcing the polynomial through every
255
+ noisy sample exactly, trading a small amount of interpolation bias
256
+ for reduced variance.
257
+
258
+ This matters in practice, not just in theory: adding more noise-scale
259
+ points to *exact* interpolation (`richardson_extrapolate`) makes
260
+ extrapolation WORSE under real statistical noise, because Lagrange
261
+ coefficients grow with point count (worse still with closely-spaced
262
+ points -- a Runge's-phenomenon-like effect). Measured directly on the
263
+ density-matrix healing experiment (`experiments/matrix_healing_zne_sweep.py`'s
264
+ setup, n=4 qubits, all 5 noise channels, 5 seeds): exact interpolation's
265
+ mean fidelity-delta dropped from +0.148 (3 points) to +0.081 (5 points,
266
+ same spacing) to -0.220 (5 points, denser spacing -- actively worse
267
+ than not correcting). A degree-2 least-squares fit fed the same extra
268
+ points instead REDUCES variance (std 0.062 -> 0.035-0.046) at
269
+ comparable or better mean delta, because the extra points are no
270
+ longer forced to satisfy an increasingly ill-conditioned exact fit.
271
+ This is why `zne_density_matrix` uses this function (degree=2) instead
272
+ of `richardson_extrapolate` by default.
273
+
274
+ Raises ValueError if fewer than `degree + 1` points are given (the fit
275
+ would be underdetermined).
276
+ """
277
+ lambdas = jnp.asarray(noise_factors, dtype=jnp.float64)
278
+ n = lambdas.shape[0]
279
+ if n < degree + 1:
280
+ raise ValueError(
281
+ f"polynomial_extrapolate needs at least degree+1={degree + 1} noise "
282
+ f"factors for a degree-{degree} fit, got {n}."
283
+ )
284
+ values_dtype = jnp.complex128 if np.iscomplexobj(np.asarray(expectation_values)) else jnp.float64
285
+ values = jnp.asarray(expectation_values, dtype=values_dtype)
286
+ return _polynomial_extrapolate_core(values, lambdas, degree)
287
+
288
+
289
+ def _polynomial_extrapolate_core(values: jnp.ndarray, lambdas: jnp.ndarray, degree: int) -> jnp.ndarray:
290
+ """`jax.jit`-traceable core of `polynomial_extrapolate` -- takes `values`
291
+ already cast to its final dtype and `lambdas` already a float64 array,
292
+ so there's no `np.iscomplexobj`/`np.asarray` call on a possibly-traced
293
+ value (which breaks tracing; `np.asarray` on a JAX tracer raises).
294
+ `degree` must be a Python int, not a traced value (`range(degree + 1)`
295
+ unrolls it at trace time) -- callers that `jax.jit` a function using
296
+ this must mark `degree` static (`static_argnames`).
297
+ """
298
+ n = lambdas.shape[0]
299
+ orig_shape = values.shape[1:]
300
+ flat = values.reshape(n, -1)
301
+ # Vandermonde design matrix: columns lambda^0, lambda^1, ..., lambda^degree.
302
+ design = jnp.stack([lambdas ** k for k in range(degree + 1)], axis=1).astype(values.dtype)
303
+ coeffs, *_ = jnp.linalg.lstsq(design, flat, rcond=None)
304
+ intercept = coeffs[0] # fitted polynomial evaluated at noise_factor=0
305
+ return intercept.reshape(orig_shape)
306
+
307
+
308
+ polynomial_extrapolate_jit = functools.partial(jax.jit, static_argnames=("degree",))(_polynomial_extrapolate_core)
309
+ """`jax.jit`-compiled entry point for `polynomial_extrapolate`, added for
310
+ consistency with every other function in this module (`richardson_extrapolate_jit`,
311
+ `zero_noise_extrapolation_jit`, `uhlmann_fidelity_jit`, `zne_density_matrix_jit`)
312
+ -- until now this was the one function whose `_core` existed (used
313
+ internally by `zne_density_matrix_jit`) but had no standalone public jit
314
+ entry point of its own.
315
+
316
+ `values` must already be cast to its final dtype (complex128 or float64)
317
+ and `lambdas` a float64 array -- this skips `polynomial_extrapolate`'s
318
+ `np.iscomplexobj` auto-detection, not traceable. `degree` is static (same
319
+ constraint as `zne_density_matrix_jit`). Verified to match
320
+ `polynomial_extrapolate` exactly."""
321
+
322
+
323
+ def bounded_exponential_extrapolate(expectation_values, noise_factors, bound: float = 1.0) -> jnp.ndarray:
324
+ """Physically bounded exponential Zero-Noise Extrapolation (Miranskyy,
325
+ Sorrenti, Thind & Gravel, arXiv:2604.24475, "Improving Zero-Noise
326
+ Extrapolation via Physically Bounded Models").
327
+
328
+ Unlike `polynomial_extrapolate`, this fits an EXPONENTIAL model --
329
+ appropriate when the expectation value decays roughly exponentially with
330
+ noise strength (the usual case for a depolarizing-type channel), not a
331
+ polynomial one -- and is Dense-Evolution's first exponential-family ZNE
332
+ model. The zero-noise value is made an explicit model parameter via the
333
+ reparametrization
334
+
335
+ E(lambda) = a + (zeta - a) * exp(-c * lambda), E(0) = zeta
336
+
337
+ and `zeta` is constrained to `[-bound, bound]` during the fit (the
338
+ physically valid range for a +-1-eigenvalue Pauli observable, `bound=1.0`
339
+ by default). An ordinary unconstrained `a + b*exp(-c*lambda)` fit has no
340
+ such guarantee and can produce a wildly out-of-range or non-finite
341
+ "zero-noise" estimate when the data is noisy or only a few points are
342
+ available -- verified directly on this project's own depolarizing-noise
343
+ ZZ-expectation setup (real Bell circuit, 30 random noise seeds, 3 noise
344
+ scales): the unconstrained fit landed outside [-1, 1] or failed to
345
+ converge in 21/30 seeds (mean absolute error over 200000, dominated by
346
+ those blow-ups), versus 0/30 for this bounded fit (mean absolute error
347
+ 0.066) -- see `Dense-Evolution-Discovery/scripts/zne_physically_bounded.py`
348
+ for the full reproduction.
349
+
350
+ Fit via SciPy's constrained L-BFGS-B, not `jax.jit`-traceable like the
351
+ rest of this module -- the optimization itself runs in plain NumPy, so
352
+ (unlike every other function here) there is no `_jit` variant. Multi-start
353
+ (`_N_MULTI_STARTS` starts: the original single fixed start
354
+ `[0, values[0], 0.5]` first, then `_N_MULTI_STARTS - 1` deterministic
355
+ random ones, fixed seed `_MULTI_START_SEED`, keeping the converged fit
356
+ with lowest loss) -- keeping the original start first means a case
357
+ where it was already the global optimum is unaffected bit-for-bit.
358
+ The non-convex 3-parameter fit from a single fixed starting point can
359
+ otherwise land in a poor local minimum -- on data generated exactly
360
+ from this function's own model, the single-start fit was off by up to
361
+ 3.9e-3, versus under 1e-5 with multi-start. Raises `RuntimeError` if every
362
+ start fails to converge (`result.success`), rather than silently
363
+ returning an unconverged `result.x[1]`.
364
+
365
+ Examples
366
+ --------
367
+ >>> from dense_evolution.mitigation import bounded_exponential_extrapolate
368
+ >>> round(float(bounded_exponential_extrapolate([0.49, 0.1, 0.1], [1.0, 2.0, 3.0])), 4)
369
+ 1.0
370
+ """
371
+ lambdas = np.asarray(noise_factors, dtype=np.float64)
372
+ values = np.asarray(expectation_values, dtype=np.float64)
373
+
374
+ def loss(theta):
375
+ a, zeta, c = theta
376
+ pred = a + (zeta - a) * np.exp(-c * lambdas)
377
+ return np.sum((values - pred) ** 2)
378
+
379
+ rng = np.random.default_rng(_MULTI_START_SEED)
380
+ starts = np.concatenate([
381
+ np.array([[0.0, float(values[0]), 0.5]]),
382
+ rng.uniform([-bound, -bound, 0.01], [bound, bound, 3.0], size=(_N_MULTI_STARTS - 1, 3)),
383
+ ])
384
+ best_result = None
385
+ for x0 in starts:
386
+ result = minimize(
387
+ loss, x0=x0, method="L-BFGS-B",
388
+ bounds=[(-bound, bound), (-bound, bound), (1e-6, None)],
389
+ options={"ftol": 1e-14, "gtol": 1e-12},
390
+ )
391
+ if not result.success:
392
+ continue
393
+ if best_result is None or result.fun < best_result.fun:
394
+ best_result = result
395
+
396
+ if best_result is None:
397
+ raise RuntimeError(
398
+ f"bounded_exponential_extrapolate: none of {_N_MULTI_STARTS} multi-start "
399
+ f"L-BFGS-B fits converged (result.success was False for every start)."
400
+ )
401
+ return jnp.float64(best_result.x[1])
402
+
403
+
404
+ def project_to_physical(rho_raw: jnp.ndarray) -> jnp.ndarray:
405
+ """Project a Hermitian, trace-1 candidate matrix onto the nearest
406
+ physical density matrix (Hermitian, trace 1, positive-semidefinite) in
407
+ 2-norm/Frobenius distance -- the same problem Smolin, Gambetta & Smith,
408
+ "Maximum Likelihood, Minimum Effort" (2012), arXiv:1106.5458, Fig. 1,
409
+ solve with a sequential eigenvalue-clipping algorithm (sort eigenvalues
410
+ descending, repeatedly zero the smallest remaining one and redistribute
411
+ its negative mass over the rest, until the least would be non-negative).
412
+
413
+ Implemented here as Euclidean projection onto the probability simplex
414
+ (Held, Wolfe & Crowder 1974; also e.g. Duchi et al. 2008) applied to the
415
+ eigenvalues -- a different, fully vectorized algorithm for the exact
416
+ same convex optimization problem (unique global minimum, so any correct
417
+ algorithm must agree). No Python-level `while`/`for` loop over
418
+ eigenvalues (the original transcription's `while i >= 0: ...` isn't
419
+ `jax.jit`-traceable, forcing a host round-trip every call) -- this
420
+ version is pure `jnp` array ops plus one dynamic index (`mus[k-1]`,
421
+ itself trace-safe), so it JIT-compiles cleanly.
422
+
423
+ Verified against the original transcription (which itself matches the
424
+ SGS paper's own worked example, eigenvalues 3/5, 1/2, 7/20, 1/10,
425
+ -11/20 -> 9/20, 7/20, 1/5, 0, 0): identical to machine precision on the
426
+ paper's example and on 30 random Hermitian trace-1 matrices (2-7 dim)
427
+ perturbed to be unphysical (max difference ~1e-15); confirmed to
428
+ actually compile and run under `jax.jit`.
429
+
430
+ `richardson_extrapolate`/`polynomial_extrapolate`'s output on a stack
431
+ of density matrices is not itself generally a valid density matrix --
432
+ extrapolation can (and in practice does) produce small negative
433
+ eigenvalues even when every input matrix was physical. This is the
434
+ correction step, meant to run *after* extrapolation, not a
435
+ general-purpose "make anything a density matrix" tool (it assumes the
436
+ input is already Hermitian and trace 1 up to this function's own
437
+ re-Hermitization step below).
438
+ """
439
+ rho_h = 0.5 * (rho_raw + jnp.conj(rho_raw).T)
440
+ eigvals, eigvecs = jnp.linalg.eigh(rho_h)
441
+ idx = jnp.argsort(eigvals)[::-1]
442
+ evals = eigvals[idx]
443
+ vecs = eigvecs[:, idx]
444
+
445
+ d = evals.shape[0]
446
+ ranks = jnp.arange(1, d + 1)
447
+ cumsum_evals = jnp.cumsum(evals)
448
+ # For each prefix length j, the shift mu_j that would make that prefix
449
+ # (plus this shift) sum to 1; the simplex-projection lemma guarantees
450
+ # `evals > mus` is a prefix-of-True mask for descending-sorted evals,
451
+ # so its count k is exactly the largest valid prefix length.
452
+ mus = (cumsum_evals - 1.0) / ranks
453
+ k = jnp.sum((evals > mus).astype(jnp.int32))
454
+ chosen_mu = mus[k - 1]
455
+
456
+ projected_evals = jnp.maximum(evals - chosen_mu, 0.0)
457
+ rho_physical = (vecs * projected_evals) @ jnp.conj(vecs).T
458
+ return jnp.asarray(rho_physical, dtype=jnp.complex128)
459
+
460
+
461
+ @jax.custom_jvp
462
+ def _eigh_degenerate_safe(A: jnp.ndarray):
463
+ """Same (eigvals, eigvecs) as `jnp.linalg.eigh(A)` -- only the backward
464
+ (gradient) rule differs, to stay finite at (near-)degenerate
465
+ eigenvalues, where JAX's own built-in `eigh` gradient rule divides by
466
+ `lambda_i - lambda_j` and returns NaN (documented upstream, e.g. JAX
467
+ issues #2311 and #8732; general treatment in Kasim, "Derivatives of
468
+ partial eigendecomposition of a real symmetric matrix for degenerate
469
+ cases", arXiv:2011.04366).
470
+
471
+ This is a practical variant of that fix: mask the `1/(lambda_i -
472
+ lambda_j)` term to 0 for near-degenerate pairs instead of letting it
473
+ blow up, rather than Kasim's fuller per-degenerate-block treatment
474
+ (which recovers a generally nonzero contribution from within the
475
+ degenerate eigenspace itself for some perturbation directions).
476
+ Sufficient here because `uhlmann_fidelity`'s only use of the
477
+ eigenVECTORS is to rebuild a matrix square root, and simply zeroing
478
+ that masked term still gives a mathematically valid (if not maximally
479
+ sharp) subgradient direction there.
480
+
481
+ This masking is NOT known to already exist elsewhere, checked directly
482
+ rather than assumed: PyTorch's current `linalg_eig_jvp`/
483
+ `linalg_eig_backward` (`torch/csrc/autograd/FunctionsManual.cpp`)
484
+ divides directly by `(L_j - L_i)` with no degenerate-pair masking at
485
+ all -- the same singularity as JAX's default, unresolved there too.
486
+ The `xitorch` library (built by Kasim himself)'s own derivation notes
487
+ (`doc/notes/deriv_symeig.rst`) state "This derivation assumes the
488
+ eigenvalues are all unique. Cases with degenerate eigenvalues are
489
+ treated differently" without giving that treatment on that page, and
490
+ its `symeig` docstring warns directly that for degenerate values "the
491
+ calculation and its gradient might be inaccurate". This appears to be
492
+ a genuinely open problem across the JAX/PyTorch/xitorch ecosystem, not
493
+ something ported from prior art.
494
+
495
+ Verified (directional-derivative/JVP comparison against symmetric-
496
+ tangent finite differences, the only comparison method that avoids
497
+ eigenvector sign/ordering convention ambiguities): exact match
498
+ (diff=0.000000) for non-degenerate, 2-fold, and 3-fold degenerate test
499
+ matrices. The 1e-8 masking threshold itself is verified, not
500
+ arbitrary: swept true (non-exact) eigenvalue gaps from 1.0 down to
501
+ 1e-10 against finite differences -- the unmasked formula tracks them
502
+ to ~1e-9 accuracy for any gap down to ~1e-7, only degrading right at
503
+ the threshold, which sits at the edge of float64-resolvable gaps for
504
+ O(1)-scale eigenvalues rather than discarding real signal."""
505
+ return jnp.linalg.eigh(A)
506
+
507
+
508
+ @_eigh_degenerate_safe.defjvp
509
+ def _eigh_degenerate_safe_jvp(primals, tangents):
510
+ A, = primals
511
+ dA, = tangents
512
+ w, v = jnp.linalg.eigh(A)
513
+ dA_sym = 0.5 * (dA + jnp.conj(dA).T)
514
+ vt_dA_v = jnp.conj(v).T @ dA_sym @ v
515
+ dw = jnp.real(jnp.diag(vt_dA_v))
516
+
517
+ denom = w[None, :] - w[:, None]
518
+ is_degenerate = jnp.abs(denom) < 1e-8
519
+ safe_denom = jnp.where(is_degenerate, 1.0, denom)
520
+ F = jnp.where(is_degenerate, 0.0, 1.0 / safe_denom)
521
+
522
+ dv = v @ (F * vt_dA_v)
523
+ return (w, v), (dw, dv)
524
+
525
+
526
+ def uhlmann_fidelity(rho_A: jnp.ndarray, rho_B: jnp.ndarray) -> float:
527
+ """Uhlmann fidelity F(rho_A, rho_B) = (Tr sqrt(sqrt(rho_A) rho_B sqrt(rho_A)))^2.
528
+
529
+ Reduces to |<psi_A|psi_B>|^2 when both inputs are pure-state density
530
+ matrices (verified directly). Validation-only: this is meant to grade
531
+ a correction against a known ideal state, never to feed into one --
532
+ passing it a target/ideal density matrix as an input to
533
+ `zne_density_matrix` or any extrapolation step would be using held-out
534
+ ground truth to guide the algorithm (oracle access), not a legitimate
535
+ error-mitigation technique. Keeping ideal-state comparison to this
536
+ function only, rather than plumbing it into the correction functions
537
+ at all, makes that boundary structural rather than a convention callers
538
+ have to remember.
539
+
540
+ Computes Tr(sqrt(inner)) as sum(sqrt(eigenvalues of inner)) instead of
541
+ reconstructing the full matrix square root (sqrt(M) has the same
542
+ eigenvectors as M and sqrt-of-eigenvalues eigenvalues, so its trace is
543
+ exactly that sum) -- skips one eigenvector reconstruction, and avoids
544
+ `matsqrt`'s `float()` cast that isn't `jax.jit`-traceable. Verified
545
+ against the previous full-reconstruction version: identical to machine
546
+ precision (~1e-16) on 30 random density-matrix pairs.
547
+
548
+ Differentiable through both arguments, including when `rho_A` has
549
+ (near-)degenerate eigenvalues (e.g. a near-pure state's noisy density
550
+ matrix, which typically has several near-zero, near-degenerate
551
+ eigenvalues) -- uses `_eigh_degenerate_safe` internally rather than
552
+ `jnp.linalg.eigh` directly, specifically to keep `jax.grad(uhlmann_fidelity, ...)`
553
+ finite in that case (see `_eigh_degenerate_safe`'s own docstring).
554
+ Forward-pass value is bit-identical to `jnp.linalg.eigh`-based
555
+ computation (same underlying `eigh` call; only the backward rule
556
+ differs), verified end-to-end against the previous implementation.
557
+
558
+ Examples
559
+ --------
560
+ >>> import numpy as np
561
+ >>> from dense_evolution.mitigation import uhlmann_fidelity
562
+ >>> rho = np.array([[1, 0], [0, 0]], dtype=complex)
563
+ >>> round(float(uhlmann_fidelity(rho, rho)), 4)
564
+ 1.0
565
+ >>> sigma = np.array([[0.5, 0], [0, 0.5]], dtype=complex)
566
+ >>> round(float(uhlmann_fidelity(rho, sigma)), 4)
567
+ 0.5
568
+ """
569
+ rho_A = jnp.asarray(rho_A, dtype=jnp.complex128)
570
+ rho_B = jnp.asarray(rho_B, dtype=jnp.complex128)
571
+ return float(_uhlmann_fidelity_core(rho_A, rho_B))
572
+
573
+
574
+ # global_depolarizing_channel, amplitude_damping_channel, and
575
+ # cosmic_ray_burst_profile moved to dense_evolution.noise -- they generate
576
+ # noise, they don't mitigate it, so this module (mitigation) was the wrong
577
+ # home. Re-exported below for backward compatibility; new code should
578
+ # import from dense_evolution.noise directly.
579
+ from ..noise import global_depolarizing_channel, amplitude_damping_channel, cosmic_ray_burst_profile
580
+
581
+
582
+ def _uhlmann_fidelity_core(rho_A: jnp.ndarray, rho_B: jnp.ndarray) -> jnp.ndarray:
583
+ """`jax.jit`-traceable core of `uhlmann_fidelity` -- both inputs already
584
+ complex128; returns a jnp scalar instead of a Python `float` (the
585
+ `float()` cast isn't traceable, same reason `_jsd_vectors_jax`/
586
+ `_jsd_vectors` are split this way in `dense_evolution.mps`).
587
+ """
588
+ def sqrt_eigvals(m):
589
+ w = jnp.linalg.eigvalsh(m)
590
+ return jnp.clip(jnp.real(w), 0.0, None)
591
+
592
+ w_A, v_A = _eigh_degenerate_safe(rho_A)
593
+ sqrt_A = (v_A * jnp.sqrt(jnp.clip(jnp.real(w_A), 0.0, None))) @ jnp.conj(v_A).T
594
+ inner = sqrt_A @ rho_B @ sqrt_A
595
+ inner_evals = sqrt_eigvals(inner)
596
+ return jnp.real(jnp.sum(jnp.sqrt(inner_evals)) ** 2)
597
+
598
+
599
+ uhlmann_fidelity_jit = jax.jit(_uhlmann_fidelity_core)
600
+ """`jax.jit`-compiled entry point for `uhlmann_fidelity`. Both `rho_A`/
601
+ `rho_B` must already be `complex128` (this skips `uhlmann_fidelity`'s
602
+ own `jnp.asarray(..., dtype=jnp.complex128)` cast, itself trace-safe, but
603
+ kept out of the core to mirror the other `_core` functions' convention).
604
+ Returns a jnp scalar, not a Python `float` -- call `float(...)` yourself
605
+ if you need one outside a jitted context. Verified to match `uhlmann_fidelity`
606
+ exactly (same underlying math, just not cast to a Python float)."""
607
+
608
+
609
+ def zne_density_matrix(rho_at_scales, noise_factors, degree: int = 2) -> jnp.ndarray:
610
+ """Zero-Noise Extrapolation for density matrices.
611
+
612
+ `rho_at_scales[i]` is a noisy density-matrix estimate (e.g. from a
613
+ Monte-Carlo/shot ensemble) at noise scale `noise_factors[i]`.
614
+ Extrapolates to zero noise via `polynomial_extrapolate` (least-squares,
615
+ complex-safe, degree=2 by default) and projects the result onto the
616
+ nearest physical density matrix (`project_to_physical`), since the raw
617
+ extrapolated matrix is not generally positive-semidefinite even when
618
+ every input was.
619
+
620
+ Uses `polynomial_extrapolate` rather than exact `richardson_extrapolate`
621
+ because, with exactly 3 noise scales (the original design point), the
622
+ two are mathematically identical -- but `polynomial_extrapolate` stays
623
+ well-behaved (reduced variance) when a caller passes MORE than 3 scales,
624
+ where exact interpolation instead gets WORSE (see
625
+ `polynomial_extrapolate`'s docstring for the measured numbers). This
626
+ makes "pass more noise-scale points" a safe thing to try rather than a
627
+ trap.
628
+
629
+ Honest findings, both against a GHZ-state ideal target and
630
+ `dense_evolution.registry.NoiseModel` noise at base_p=0.05, scales
631
+ 1x/2x/3x unless noted, `uhlmann_fidelity` against the true ideal state
632
+ used only to grade the result -- never as input to any step above:
633
+
634
+ - `experiments/matrix_healing_zne.py`: 2-qubit Bell state, depolarizing
635
+ noise, K=200-trajectory estimate per scale, averaged over 4 seeds --
636
+ raw fidelity ~0.865, corrected ~0.947 (+0.08), positive on every seed
637
+ tested.
638
+ - `experiments/matrix_healing_zne_sweep.py`: 2-5 qubits x all 5
639
+ `NoiseModel` channels (depolarizing, bitflip, phaseflip,
640
+ amplitude_damping, combined) x 5 seeds, K=400 trajectories per scale
641
+ (100 runs total) -- **96/100 positive, mean delta +0.12**, and every
642
+ single (qubit count, noise channel) combination is net positive on
643
+ average, growing to +0.20-0.25 at 5 qubits for depolarizing/bitflip.
644
+ The 4 remaining negative runs are small (worst -0.02) and consistent
645
+ with residual Monte Carlo noise, not a systematic failure mode.
646
+ (These specific numbers were measured with exact 3-point
647
+ interpolation, which is identical to this function's degree=2
648
+ default at 3 points -- unaffected by the switch.)
649
+ - An earlier draft of this sweep (K=150, 3 seeds) had reported
650
+ phaseflip/amplitude_damping as "unreliable" -- re-investigated rather
651
+ than trusted, and confirmed to be a Monte Carlo undersampling
652
+ artifact (extrapolation coefficients amplify input noise; an
653
+ undersampled estimate makes the *corrected* result noisy even when
654
+ the correction itself is sound), not a real limitation.
655
+ - More noise-scale points, SAME total measurement budget (the fair
656
+ comparison -- splitting a fixed number of trajectories across more
657
+ points, not spending more): 3 points x K=400 (1200 total) vs. 5
658
+ points x K=240 (1200 total) vs. 7 points x K=171 (~1200 total),
659
+ n=4 qubits, all 5 noise channels, 5 seeds. 5 points matches or
660
+ slightly beats the 3-point mean delta (+0.150 vs +0.148) with 19%
661
+ lower variance (std 0.050 vs 0.062) -- a real, free improvement at
662
+ the same experimental cost, not an artifact of spending more. 7
663
+ points trades a little mean (+0.132) for still-lower variance (std
664
+ 0.043, 30% below baseline) -- a genuine tradeoff point, useful when
665
+ reliability matters more than average performance. At exactly 3
666
+ points this function is mathematically identical to
667
+ `richardson_extrapolate` (verified to 1e-12) -- there is no free
668
+ lunch there, the gain only appears once more points are used.
669
+ - More noise-scale points, FIXED K per point instead (spending more
670
+ total measurement, K=400 at every point count): with exact
671
+ interpolation this makes things worse, not better (mean delta drops
672
+ from +0.148 at 3 points to +0.081 at 5, and to -0.220 at 5
673
+ closely-spaced points -- see `polynomial_extrapolate`'s docstring);
674
+ with this function's degree=2 default it instead stays comparable
675
+ in mean with lower variance (std 0.062 -> 0.035-0.046) -- confirms
676
+ the safety property holds even when *not* holding budget fixed, on
677
+ top of the fixed-budget gain above.
678
+
679
+ Practical implication for callers regardless of `degree`: correction
680
+ quality still depends on `rho_at_scales` being a reasonably low-noise
681
+ estimate to begin with (large enough K, or equivalent) -- an
682
+ extrapolation fit through pure noise cannot recover signal that isn't
683
+ there. `degree` trades bias for variance: higher degree fits the true
684
+ curve's shape more closely (less bias) but is more sensitive to
685
+ per-point noise (more variance); degree=2 was chosen empirically as the
686
+ best tested tradeoff, not a theoretical optimum for every regime.
687
+
688
+ Do not pass a target/ideal density matrix into this function or use one
689
+ to pick among candidate corrections -- see `uhlmann_fidelity`'s
690
+ docstring for why.
691
+ """
692
+ extrapolated = polynomial_extrapolate(rho_at_scales, noise_factors, degree=degree)
693
+ return project_to_physical(extrapolated)
694
+
695
+
696
+ def _zne_density_matrix_core(rho_at_scales: jnp.ndarray, noise_factors: jnp.ndarray,
697
+ degree: int) -> jnp.ndarray:
698
+ """`jax.jit`-traceable core of `zne_density_matrix` -- `rho_at_scales`
699
+ must already be complex128, `noise_factors` a float64 array; `degree`
700
+ must be a Python int (unrolled at trace time by
701
+ `_polynomial_extrapolate_core`, same constraint as there).
702
+ """
703
+ extrapolated = _polynomial_extrapolate_core(rho_at_scales, noise_factors, degree)
704
+ return project_to_physical(extrapolated)
705
+
706
+
707
+ zne_density_matrix_jit = functools.partial(jax.jit, static_argnames=("degree",))(_zne_density_matrix_core)
708
+ """`jax.jit`-compiled entry point for `zne_density_matrix`, for callers
709
+ inside a jitted pipeline (e.g. `jax.lax.scan` in `MPSSimulator.run_circuit_jit`)
710
+ who don't want a host round-trip every call -- `zne_density_matrix` itself
711
+ stays eager (unchanged) for one-off/interactive use, where jit compilation
712
+ overhead isn't worth paying for a single call.
713
+
714
+ `degree` is a static argument (must be a Python int, not a traced value --
715
+ pass it positionally or by keyword the same way every call, since JAX
716
+ recompiles per distinct static value). `rho_at_scales` must already be
717
+ `complex128` and `noise_factors` a plain float array/sequence -- unlike
718
+ `zne_density_matrix`, this skips the `np.iscomplexobj` dtype auto-detection
719
+ (not traceable) and always assumes complex input, which is the only case
720
+ that makes sense for density matrices.
721
+
722
+ Measured speedup is real but size- and call-pattern-dependent (`project_to_physical`
723
+ alone measured 2x-22x across 2x2 to 32x32 matrices when jitted vs. the
724
+ previous non-jittable version) -- benchmark your own use case rather than
725
+ assuming a fixed number; the benefit only appears once compiled and called
726
+ repeatedly, a single one-off call pays the compilation cost first.
727
+ """
728
+
729
+
730
+ def _js_divergence(p: jnp.ndarray, q: jnp.ndarray, eps: float = 1e-12) -> jnp.ndarray:
731
+ """Standard Jensen-Shannon divergence (natural log, bounded in
732
+ [0, ln 2]) between two probability vectors -- NOT the same quantity
733
+ as `dense_evolution.mps._jsd_vectors_jax` (a different, MPS-specific
734
+ "adaptive Jensen-Shannon Distance": base-2 log, an extra
735
+ log10(dim)/2 dimensional scaling factor, and a final square root,
736
+ purpose-built for sizing truncated bond dimensions). Both are
737
+ legitimately named "JSD" for their own purposes; they are not
738
+ interchangeable, and this one is the plain textbook formula, chosen
739
+ here to exactly match the already-validated implementation in
740
+ Dense-Evolution-Discovery's channel_order_noncommutativity.py."""
741
+ p, q = p + eps, q + eps
742
+ p, q = p / jnp.sum(p), q / jnp.sum(q)
743
+ m = 0.5 * (p + q)
744
+ kl = lambda a, b: jnp.sum(a * jnp.log(a / b))
745
+ return 0.5 * kl(p, m) + 0.5 * kl(q, m)
746
+
747
+
748
+ def jsd_predictive_zne_density_matrix(rho_at_scales, noise_factors) -> jnp.ndarray:
749
+ """Density-matrix ZNE with a Jensen-Shannon-divergence-informed
750
+ coefficient nudge, for noise whose scale-to-output-distribution
751
+ relationship isn't perfectly smooth (the assumption plain
752
+ 3-point Richardson extrapolation, which `zne_density_matrix`
753
+ defaults toward at exactly 3 scales, relies on).
754
+
755
+ Motivated by and validated in Dense-Evolution-Discovery's
756
+ scripts/photonic_predictive_zne.py -- prototyped there first for
757
+ photon-loss noise (a photonic-relevant channel: photon loss on a
758
+ dual-rail-encoded qubit IS this library's `amplitude_damping`
759
+ channel), per this project's cross-repo promotion pattern. Grounded
760
+ in real literature: Mills & Mezher, "Mitigating photon loss in
761
+ linear optical quantum circuits" (arXiv:2405.02278), find plain
762
+ scalar ZNE does not beat postselection for discrete-variable photon
763
+ loss -- reproduced directly there (scalar ZNE went unphysical,
764
+ fidelity > 1.0, at 14/16 swept points). `zne_density_matrix` avoids
765
+ that failure mode by construction (`project_to_physical`); this
766
+ function asks whether a further, data-driven adaptive correction on
767
+ top of it can do even better.
768
+
769
+ Signal: Jensen-Shannon divergence (`_js_divergence`, the standard
770
+ formula -- see its own docstring for how this differs from
771
+ `dense_evolution.mps`'s unrelated "adaptive JSD") between the
772
+ measurement-probability distributions (density-matrix diagonals) at
773
+ consecutive noise scales. Needs no external calibration or oracle
774
+ access to an ideal/target state -- unlike naively reusing
775
+ `calculate_delta_preemp` with an externally-supplied signal (tried
776
+ first in the Discovery prototype; found to have a negligible effect
777
+ by construction, since that formula's fixed nudge constants 0.01/
778
+ 0.02 were tuned for a differently-scaled use case elsewhere in this
779
+ module).
780
+
781
+ `nonlinearity = (jsd_23 - jsd_12) / (jsd_23 + jsd_12 + eps)`
782
+ (bounded in [-1, 1]) measures how consistently the JSD grows between
783
+ consecutive scales -- near 0 when the noise-scale -> output-
784
+ distribution map is locally well-behaved (Richardson's implicit
785
+ assumption holding), away from 0 when it isn't. RECTIFIED: the
786
+ coefficient nudge is applied only when `nonlinearity > 0` -- an
787
+ unrectified first version, applying the nudge for both signs,
788
+ helped in only 5/16 points on a real run despite the signal itself
789
+ being significantly correlated with success (Pearson r=+0.533,
790
+ p=0.0334); the fix was clipping to the regime the signal was shown
791
+ to work in, not discarding the signal. When `nonlinearity <= 0`,
792
+ this reduces EXACTLY to `zne_density_matrix` at 3 equally-spaced
793
+ scales (verified: max deviation ~1e-8, floating-point noise) --
794
+ zero risk in that regime, by construction.
795
+
796
+ Verified on a real, seed-diverse sample (72 points: 12 photon-loss
797
+ rates x 6 independent seeds, K=200 trajectories each) before being
798
+ promoted here, not just the small sample that first suggested it:
799
+ among 46 points where the mechanism actually activates
800
+ (nonlinearity > 0.01), 76.1% (35/46) improve over plain
801
+ `zne_density_matrix`, mean fidelity gain +0.0055, one-sample
802
+ t-test against zero p=0.0003 -- and positive in 6/6 independent
803
+ seeds tested (not one lucky seed driving the result). The win rate
804
+ and effect size were LARGER on the big sample than the small one
805
+ that first suggested it, the opposite of the usual small-sample-
806
+ regresses-to-null pattern -- checked directly rather than assumed
807
+ either way before trusting it.
808
+
809
+ Only defined for exactly 3 equally-spaced noise factors (1x, 2x,
810
+ 3x), same restriction as `zero_noise_extrapolation`'s own healing-
811
+ adapted branch and for the same reason: the underlying Lagrange
812
+ coefficients (3, -3, 1) this nudges are specific to that spacing,
813
+ not a general n-point formula."""
814
+ rho_at_scales = jnp.asarray(rho_at_scales, dtype=jnp.complex128)
815
+ noise_factors = jnp.asarray(noise_factors, dtype=jnp.float64)
816
+ if noise_factors.shape[0] != 3:
817
+ raise NotImplementedError(
818
+ "jsd_predictive_zne_density_matrix is only defined for exactly 3 noise "
819
+ "factors; call zne_density_matrix(...) directly for the plain N-point case."
820
+ )
821
+ return _jsd_predictive_zne_density_matrix_core(rho_at_scales)
822
+
823
+
824
+ def _jsd_predictive_zne_density_matrix_core(rho_at_scales: jnp.ndarray, nudge_scale: float = 0.5) -> jnp.ndarray:
825
+ """`jax.jit`-traceable core of `jsd_predictive_zne_density_matrix` --
826
+ `rho_at_scales` already complex128. See the public wrapper's
827
+ docstring for the method and its validation."""
828
+ probs = jnp.real(jnp.diagonal(rho_at_scales, axis1=-2, axis2=-1))
829
+ jsd_12 = _js_divergence(probs[0], probs[1])
830
+ jsd_23 = _js_divergence(probs[1], probs[2])
831
+ nonlinearity = (jsd_23 - jsd_12) / (jsd_23 + jsd_12 + 1e-12)
832
+ rectified = jnp.maximum(nonlinearity, 0.0)
833
+
834
+ e_l1, e_l2, e_l3 = rho_at_scales[0], rho_at_scales[1], rho_at_scales[2]
835
+ c1 = 3.0 - nudge_scale * rectified
836
+ c2 = -3.0 + 2.0 * nudge_scale * rectified
837
+ c3 = 1.0 - nudge_scale * rectified
838
+ extrapolated = (c1 * e_l1 + c2 * e_l2 + c3 * e_l3) / (c1 + c2 + c3)
839
+ return project_to_physical(extrapolated)
840
+
841
+
842
+ def _coherence_l1(rho: jnp.ndarray) -> jnp.ndarray:
843
+ """l1-norm of coherence (Baumgratz, Cramer & Plenio, Phys. Rev. Lett.
844
+ 113, 140401, 2014): the sum of the magnitudes of every off-diagonal
845
+ density-matrix entry. A standard, basis-dependent measure of how much
846
+ quantum coherence a state carries in the computational basis --
847
+ unlike `_js_divergence`, which only ever sees the diagonal
848
+ (populations), this is sensitive to exactly what dephasing destroys."""
849
+ n = rho.shape[0]
850
+ return jnp.sum(jnp.abs(rho) * (1.0 - jnp.eye(n, dtype=rho.dtype)))
851
+
852
+
853
+ def coherence_predictive_zne_density_matrix(rho_at_scales, noise_factors) -> jnp.ndarray:
854
+ """Density-matrix ZNE with a coherence-informed coefficient nudge --
855
+ the same adaptive-nonlinearity mechanism as
856
+ `jsd_predictive_zne_density_matrix`, but signaled by the l1-norm of
857
+ coherence (`_coherence_l1`) instead of the Jensen-Shannon divergence
858
+ of the diagonal populations.
859
+
860
+ Motivation: `jsd_predictive_zne_density_matrix`'s signal is the
861
+ density-matrix diagonal only. Any purely dephasing-type noise
862
+ (phase-flip, or a coherent Z-axis over-rotation) is diagonal in the
863
+ computational basis -- it moves phase, never populations -- so that
864
+ signal is blind to it BY CONSTRUCTION, not merely weak: verified
865
+ directly in Dense-Evolution-Discovery's
866
+ scripts/jsd_zne_noise_generalization.py, the fidelity delta from the
867
+ classical-JSD nudge is exactly 0.0 at every tested phase-flip noise
868
+ strength and every tested coherent-rotation angle. A quantum-JSD
869
+ variant (von Neumann entropy of the full density matrix instead of
870
+ Shannon entropy of the diagonal) was tried there too and rejected: it
871
+ weakens the already-working amplitude-damping/combined-noise case
872
+ without fixing the coherent-error case, since a smooth deterministic
873
+ function of the noise-scale factor has `jsd_12~=jsd_23` regardless of
874
+ which divergence measures it -- the nonlinearity trigger this whole
875
+ family of methods relies on is structurally near-zero there no
876
+ matter the signal.
877
+
878
+ Validated scope, checked directly rather than assumed universal: real
879
+ effect on phase-flip/dephasing-dominated noise for `base_p<=0.10`;
880
+ NOT validated (and not claimed) for amplitude-damping-dominated noise
881
+ (use `jsd_predictive_zne_density_matrix` there instead) or for
882
+ coherent/deterministic errors (out of reach for this entire family of
883
+ methods, not just this signal -- see above).
884
+
885
+ Verified at 200 independent seeds on phase-flip noise (GHZ(4),
886
+ `base_p=0.05`, K=150 trajectories/scale) before promotion: the nudge
887
+ activates (rectified nonlinearity > 0) on 63/200 seeds (31.5%) --
888
+ among those, 63/63 improve over plain `zne_density_matrix`, mean
889
+ fidelity gain +0.014892, one-sample t-test against zero
890
+ p=1.07e-08, a 20000-resample permutation test finding no resample
891
+ matching or exceeding the observed effect (p<0.00005). Confirmed
892
+ across a noise-level sweep (100 seeds/level, GHZ(4)): significant by
893
+ both tests, 100% win rate among active points, at base_p in (0.03,
894
+ 0.05, 0.08, 0.10); NOT significant at base_p=0.15 (p=0.288 t-test,
895
+ p=0.301 permutation) -- a real, honest upper boundary, not a
896
+ universal effect at any noise strength. Confirmed on a second circuit
897
+ family (hardware-efficient VQE-style ansatz, 2 layers) at
898
+ `base_p=0.05`: 69/150 active (46%, a higher activation rate than
899
+ GHZ), 67/69 positive, p=4.4e-06 -- the effect is not GHZ-specific.
900
+ When inactive, reduces EXACTLY to `zne_density_matrix` at 3
901
+ equally-spaced scales (verified: max deviation ~1e-8, floating-point
902
+ noise) -- zero risk in that regime, by construction, the same safety
903
+ property `jsd_predictive_zne_density_matrix` has.
904
+
905
+ Only defined for exactly 3 equally-spaced noise factors (1x, 2x, 3x),
906
+ same restriction and reason as `jsd_predictive_zne_density_matrix`:
907
+ the Lagrange coefficients (3, -3, 1) this nudges are specific to that
908
+ spacing."""
909
+ rho_at_scales = jnp.asarray(rho_at_scales, dtype=jnp.complex128)
910
+ noise_factors = jnp.asarray(noise_factors, dtype=jnp.float64)
911
+ if noise_factors.shape[0] != 3:
912
+ raise NotImplementedError(
913
+ "coherence_predictive_zne_density_matrix is only defined for exactly 3 noise "
914
+ "factors; call zne_density_matrix(...) directly for the plain N-point case."
915
+ )
916
+ return _coherence_predictive_zne_density_matrix_core(rho_at_scales)
917
+
918
+
919
+ def _coherence_predictive_zne_density_matrix_core(rho_at_scales: jnp.ndarray, nudge_scale: float = 0.5) -> jnp.ndarray:
920
+ """`jax.jit`-traceable core of `coherence_predictive_zne_density_matrix`
921
+ -- `rho_at_scales` already complex128. See the public wrapper's
922
+ docstring for the method and its validation."""
923
+ c1_, c2_, c3_ = (_coherence_l1(rho_at_scales[i]) for i in range(3))
924
+ gap_12 = jnp.abs(c1_ - c2_)
925
+ gap_23 = jnp.abs(c2_ - c3_)
926
+ nonlinearity = (gap_23 - gap_12) / (gap_23 + gap_12 + 1e-12)
927
+ rectified = jnp.maximum(nonlinearity, 0.0)
928
+
929
+ e_l1, e_l2, e_l3 = rho_at_scales[0], rho_at_scales[1], rho_at_scales[2]
930
+ c1 = 3.0 - nudge_scale * rectified
931
+ c2 = -3.0 + 2.0 * nudge_scale * rectified
932
+ c3 = 1.0 - nudge_scale * rectified
933
+ extrapolated = (c1 * e_l1 + c2 * e_l2 + c3 * e_l3) / (c1 + c2 + c3)
934
+ return project_to_physical(extrapolated)
935
+
936
+
937
+ def classically_augmented_zne_phaseflip(rho_at_scales_measured, noise_factors, rho_ideal, base_p, degree: int = 2) -> jnp.ndarray:
938
+ """Classically Augmented ZNE (Scheiber et al., arXiv:2607.25746) for
939
+ phaseflip noise: the highest-noise Richardson/polynomial-extrapolation
940
+ nodes -- the ones contributing most to sampling variance -- are
941
+ replaced by `phaseflip_channel_exact`'s zero-sampling-variance exact
942
+ channel instead of a Monte-Carlo-sampled density matrix, then combined
943
+ exactly as `zne_density_matrix` already does (same extrapolation
944
+ coefficients; only the source of the high-noise inputs changes).
945
+
946
+ `rho_at_scales_measured[i]` must be the measured/sampled density
947
+ matrix at `noise_factors[i]` for `i < len(rho_at_scales_measured)`
948
+ (the low-noise nodes, kept as real measurements); every remaining
949
+ factor in `noise_factors` is filled in with the exact channel computed
950
+ from `rho_ideal` and `base_p`. `rho_ideal` and `base_p` are required
951
+ inputs, not optional, because the exact channel needs the noise-free
952
+ state to condition on -- unlike every other function in this module,
953
+ this one is not usable from noisy measurements alone.
954
+
955
+ Honest, verified scope (GHZ(3), phaseflip, 150 trials/measured node,
956
+ 60 seeds each): with exactly 3 total noise factors and only the single
957
+ highest one replaced, no measurable benefit (variance ratio 0.99x,
958
+ i.e. no effect within noise) -- too little room for the exact node to
959
+ matter against only 2 remaining measured ones. With 5 noise factors
960
+ (1x-5x, `base_p=0.03`) and the top 3 replaced by the exact channel,
961
+ variance drops by a real, measured 1.30x versus plain `zne_density_matrix`
962
+ using 5 fully-measured nodes at the same per-node trial budget. This is
963
+ the naive-allocation regime, not the paper's own optimal importance-
964
+ sampling allocation (their Eq. 9) -- that allocation is not implemented
965
+ here, so the exponential variance reduction the paper reports under it
966
+ is not claimed or expected from this function as shipped.
967
+
968
+ Only implemented for phaseflip noise, since `phaseflip_channel_exact`
969
+ is the only exact density-matrix channel currently available that
970
+ matches a `NoiseModel` statevector model exactly (see that function's
971
+ docstring) -- extending this to other noise models needs their own
972
+ exact channel first."""
973
+ from ..noise import phaseflip_channel_exact
974
+
975
+ rho_at_scales_measured = jnp.asarray(rho_at_scales_measured, dtype=jnp.complex128)
976
+ noise_factors = jnp.asarray(noise_factors, dtype=jnp.float64)
977
+ rho_ideal = jnp.asarray(rho_ideal, dtype=jnp.complex128)
978
+ cutoff = rho_at_scales_measured.shape[0]
979
+ if cutoff >= noise_factors.shape[0]:
980
+ raise ValueError(
981
+ f"rho_at_scales_measured has {cutoff} entries but noise_factors only has "
982
+ f"{noise_factors.shape[0]} -- at least one factor must be left for the exact "
983
+ "classical node, or there is nothing to augment."
984
+ )
985
+ classical_rhos = jnp.stack([
986
+ phaseflip_channel_exact(rho_ideal, jnp.minimum(base_p * f, 1.0))
987
+ for f in noise_factors[cutoff:]
988
+ ])
989
+ rho_at_scales = jnp.concatenate([rho_at_scales_measured, classical_rhos], axis=0)
990
+ return zne_density_matrix(rho_at_scales, noise_factors, degree=degree)