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,432 @@
|
|
|
1
|
+
import shutil
|
|
2
|
+
import tempfile
|
|
3
|
+
from typing import List, Optional
|
|
4
|
+
|
|
5
|
+
import numpy as np
|
|
6
|
+
|
|
7
|
+
from ._engine_imports import DenseSVSimulator
|
|
8
|
+
from .guard import HAS_JAX, MemoryPressureError, SafeMemoryGuard
|
|
9
|
+
from .geometry import MemoryChunker
|
|
10
|
+
from .circuit_chunker import CircuitChunker
|
|
11
|
+
from .kernels import (
|
|
12
|
+
_build_multi_chunk_runner, _build_distributed_chunk_runner,
|
|
13
|
+
_compile_multi_chunk_ops,
|
|
14
|
+
)
|
|
15
|
+
from .disk_overflow import run_disk_overflow_circuit
|
|
16
|
+
|
|
17
|
+
__all__ = ["Chunk", "chunk1", "chunk2", "Chunk2Incrociato"]
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
# ─────────────────────────────────────────────────────────────────────────────────
|
|
21
|
+
# Chunk (chunk2 / Chunk2Incrociato)
|
|
22
|
+
# ─────────────────────────────────────────────────────────────────────────────────
|
|
23
|
+
|
|
24
|
+
class Chunk:
|
|
25
|
+
"""
|
|
26
|
+
Anti-OOM wrapper for large-qubit simulation.
|
|
27
|
+
|
|
28
|
+
Does NOT subclass DenseSVSimulator directly — the parent __init__ allocates
|
|
29
|
+
2**n_qubits elements immediately (17 GB for 30 qubits).
|
|
30
|
+
|
|
31
|
+
For n_qubits <= chunk_size_bits (the RAM-safe budget): a single inner
|
|
32
|
+
simulator is allocated and the logical qubit count is stored separately.
|
|
33
|
+
|
|
34
|
+
For n_qubits > chunk_size_bits: num_chunks separate chunk_size_bits-qubit
|
|
35
|
+
simulators are held in RAM simultaneously (see _dispatch_multi) — as
|
|
36
|
+
many chunks as actually fit in RAM at once, checked up front via
|
|
37
|
+
SafeMemoryGuard.check_allocation before anything is allocated.
|
|
38
|
+
Benchmark attributes (num_chunks, chunk_size_bits, dtype) are
|
|
39
|
+
forwarded transparently from the embedded MemoryChunker.
|
|
40
|
+
|
|
41
|
+
Past that RAM ceiling, allow_disk_overflow=True (default False) falls
|
|
42
|
+
back to disk-backed storage instead of raising MemoryPressureError --
|
|
43
|
+
see dense_evolution/backends/chunk/disk_overflow.py (Pednault et al.
|
|
44
|
+
2019, arXiv:1910.09534) for the phased execution this uses, and
|
|
45
|
+
docs/api/chunk.md for a real verified demo and the speed trade-off
|
|
46
|
+
(never more than 1-2 chunks materialized in RAM at once, but every
|
|
47
|
+
gate now pays disk I/O -- v1 is correctness-first, not fast).
|
|
48
|
+
|
|
49
|
+
A SafeMemoryGuard fires before any simulator is instantiated
|
|
50
|
+
(pre-allocation check) and is also embedded in CircuitChunker for
|
|
51
|
+
per-slice protection during execution (n_qubits <= chunk_size_bits path).
|
|
52
|
+
|
|
53
|
+
Parameters
|
|
54
|
+
----------
|
|
55
|
+
n_qubits : logical qubit count of the target circuit
|
|
56
|
+
chunk_size_gates : gate-slice size for JIT compilation (default 500)
|
|
57
|
+
memory_threshold : free-RAM fraction below which execution is blocked
|
|
58
|
+
(default 0.15 = 15%)
|
|
59
|
+
use_float32 : forwarded to DenseSVSimulator
|
|
60
|
+
allow_disk_overflow : fall back to disk-backed chunks instead of
|
|
61
|
+
raising MemoryPressureError when num_chunks
|
|
62
|
+
chunks don't fit in RAM at once (default False)
|
|
63
|
+
disk_dir : directory for the overflow .npy files (default:
|
|
64
|
+
a fresh tempfile.mkdtemp(), removed by close())
|
|
65
|
+
"""
|
|
66
|
+
|
|
67
|
+
def __init__(
|
|
68
|
+
self,
|
|
69
|
+
n_qubits: int,
|
|
70
|
+
chunk_size_gates: int = 500,
|
|
71
|
+
memory_threshold: float = 0.15,
|
|
72
|
+
use_float32: bool = False,
|
|
73
|
+
allow_disk_overflow: bool = False,
|
|
74
|
+
disk_dir: Optional[str] = None,
|
|
75
|
+
):
|
|
76
|
+
# 1. Geometry — purely RAM-based, no JAX allocation yet
|
|
77
|
+
self._mem_chunker = MemoryChunker(n_qubits)
|
|
78
|
+
self._guard = SafeMemoryGuard(threshold_pct=memory_threshold)
|
|
79
|
+
|
|
80
|
+
# 2. Logical qubit count (for circuit parsing)
|
|
81
|
+
self.n = n_qubits
|
|
82
|
+
self.chunk_size_gates = chunk_size_gates
|
|
83
|
+
self._m = n_qubits - self._mem_chunker.chunk_size_bits # chunk-select qubit count (0 if num_chunks==1)
|
|
84
|
+
|
|
85
|
+
self._chunk_paths = None
|
|
86
|
+
self._disk_dir = None
|
|
87
|
+
self._owns_disk_dir = False
|
|
88
|
+
|
|
89
|
+
if self._mem_chunker.num_chunks == 1:
|
|
90
|
+
# 3a. Pre-allocation RAM check — block here rather than inside JAX
|
|
91
|
+
safe_q = min(n_qubits, self._mem_chunker.chunk_size_bits)
|
|
92
|
+
self._guard.check(f"Chunk.__init__ — allocating {safe_q}-qubit simulator")
|
|
93
|
+
|
|
94
|
+
# 4a. Physical simulator sized to what RAM can actually hold
|
|
95
|
+
self._inner_sim = DenseSVSimulator(
|
|
96
|
+
safe_q,
|
|
97
|
+
use_float32=use_float32,
|
|
98
|
+
)
|
|
99
|
+
self._chunk_sims = None
|
|
100
|
+
self._multi_chunk_runner = None
|
|
101
|
+
|
|
102
|
+
# 5a. Circuit chunker wired to the physical simulator, with same threshold
|
|
103
|
+
self._circuit_chunker = CircuitChunker(
|
|
104
|
+
simulator_instance=self._inner_sim,
|
|
105
|
+
memory_threshold=memory_threshold,
|
|
106
|
+
)
|
|
107
|
+
else:
|
|
108
|
+
# 3b. Sized pre-allocation check: num_chunks chunk-sized simulators
|
|
109
|
+
# held in RAM at once, plus ~2 chunks of headroom for the temporary
|
|
110
|
+
# arrays the cross-chunk gate-mixing math allocates at its peak.
|
|
111
|
+
num_chunks = self._mem_chunker.num_chunks
|
|
112
|
+
per_chunk_mb = self._mem_chunker.memory_mb()
|
|
113
|
+
required_mb = (num_chunks + 2) * per_chunk_mb
|
|
114
|
+
try:
|
|
115
|
+
self._guard.check_allocation(
|
|
116
|
+
required_mb,
|
|
117
|
+
f"Chunk.__init__ — allocating {num_chunks} chunks of "
|
|
118
|
+
f"{self._mem_chunker.chunk_size_bits} qubits each",
|
|
119
|
+
)
|
|
120
|
+
except MemoryPressureError:
|
|
121
|
+
if not allow_disk_overflow:
|
|
122
|
+
raise
|
|
123
|
+
self._init_disk_overflow(num_chunks, disk_dir)
|
|
124
|
+
return
|
|
125
|
+
|
|
126
|
+
# 4b. num_chunks independent chunk-sized simulators. Each one's own
|
|
127
|
+
# __init__ resets it to |0...0>: only chunk 0 should carry the
|
|
128
|
+
# amplitude-1 seed for the LOGICAL |0...0>, the rest must start
|
|
129
|
+
# at all-zero (direct .sv assignment — set_state/set_initial_state
|
|
130
|
+
# reject zero-norm vectors by design).
|
|
131
|
+
self._chunk_sims = [
|
|
132
|
+
DenseSVSimulator(self._mem_chunker.chunk_size_bits, use_float32=use_float32)
|
|
133
|
+
for _ in range(num_chunks)
|
|
134
|
+
]
|
|
135
|
+
for sim in self._chunk_sims[1:]:
|
|
136
|
+
sim.sv = sim.xp.zeros(self._mem_chunker.chunk_dim, dtype=sim.dtype)
|
|
137
|
+
|
|
138
|
+
self._inner_sim = None
|
|
139
|
+
self._circuit_chunker = None
|
|
140
|
+
|
|
141
|
+
# 6b. JIT runner for the multi-chunk dispatch — built once here
|
|
142
|
+
# since num_chunks/m/chunk_size_bits are fixed for this
|
|
143
|
+
# instance's whole lifetime, reused by every run_chunk() call.
|
|
144
|
+
self._multi_chunk_runner = _build_multi_chunk_runner(
|
|
145
|
+
num_chunks, self._m, self._mem_chunker.chunk_size_bits)
|
|
146
|
+
|
|
147
|
+
# 7b. Distributed (multi-device) runner — built lazily on first
|
|
148
|
+
# run_chunk_distributed() call, not here: it requires
|
|
149
|
+
# jax.device_count() >= num_chunks, which most single-process
|
|
150
|
+
# uses of Chunk will never need or satisfy.
|
|
151
|
+
self._distributed_runner = None
|
|
152
|
+
self._distributed_mesh = None
|
|
153
|
+
|
|
154
|
+
def _init_disk_overflow(self, num_chunks: int, disk_dir: Optional[str]) -> None:
|
|
155
|
+
"""Fallback storage for allow_disk_overflow=True when num_chunks
|
|
156
|
+
chunks don't fit in RAM at once -- see disk_overflow.py. Each
|
|
157
|
+
chunk becomes its own .npy file instead of a live DenseSVSimulator;
|
|
158
|
+
chunk 0 seeded to the logical |0...0>, the rest zero, same
|
|
159
|
+
convention _chunk_sims uses in the in-RAM path."""
|
|
160
|
+
self._inner_sim = None
|
|
161
|
+
self._circuit_chunker = None
|
|
162
|
+
self._chunk_sims = None
|
|
163
|
+
self._multi_chunk_runner = None
|
|
164
|
+
self._distributed_runner = None
|
|
165
|
+
self._distributed_mesh = None
|
|
166
|
+
|
|
167
|
+
self._owns_disk_dir = disk_dir is None
|
|
168
|
+
self._disk_dir = disk_dir or tempfile.mkdtemp(prefix="dense_evolution_chunk_")
|
|
169
|
+
dtype = self._mem_chunker.dtype
|
|
170
|
+
chunk_dim = self._mem_chunker.chunk_dim
|
|
171
|
+
paths = []
|
|
172
|
+
for i in range(num_chunks):
|
|
173
|
+
arr = np.zeros(chunk_dim, dtype=dtype)
|
|
174
|
+
if i == 0:
|
|
175
|
+
arr[0] = 1.0
|
|
176
|
+
path = f"{self._disk_dir}/chunk_{i}.npy"
|
|
177
|
+
np.save(path, arr)
|
|
178
|
+
paths.append(path)
|
|
179
|
+
self._chunk_paths = paths
|
|
180
|
+
|
|
181
|
+
def close(self) -> None:
|
|
182
|
+
"""Removes the disk-overflow directory, if this Chunk created one
|
|
183
|
+
(allow_disk_overflow=True with no explicit disk_dir). Safe to call
|
|
184
|
+
even if disk overflow was never used."""
|
|
185
|
+
if self._owns_disk_dir and self._disk_dir is not None:
|
|
186
|
+
shutil.rmtree(self._disk_dir, ignore_errors=True)
|
|
187
|
+
self._disk_dir = None
|
|
188
|
+
|
|
189
|
+
def __del__(self):
|
|
190
|
+
try:
|
|
191
|
+
self.close()
|
|
192
|
+
except Exception:
|
|
193
|
+
pass
|
|
194
|
+
|
|
195
|
+
# ── Benchmark-facing attribute forwarding ─────────────────────
|
|
196
|
+
|
|
197
|
+
@property
|
|
198
|
+
def num_chunks(self) -> int:
|
|
199
|
+
return self._mem_chunker.num_chunks
|
|
200
|
+
|
|
201
|
+
@property
|
|
202
|
+
def chunk_size_bits(self) -> int:
|
|
203
|
+
return self._mem_chunker.chunk_size_bits
|
|
204
|
+
|
|
205
|
+
@property
|
|
206
|
+
def chunk_dim(self) -> int:
|
|
207
|
+
return self._mem_chunker.chunk_dim
|
|
208
|
+
|
|
209
|
+
@property
|
|
210
|
+
def dtype(self):
|
|
211
|
+
return self._mem_chunker.dtype
|
|
212
|
+
|
|
213
|
+
@property
|
|
214
|
+
def memory_geometry(self) -> MemoryChunker:
|
|
215
|
+
return self._mem_chunker
|
|
216
|
+
|
|
217
|
+
# ── Simulator-facing forwarding ─────────────────────────────
|
|
218
|
+
|
|
219
|
+
@property
|
|
220
|
+
def sv(self):
|
|
221
|
+
"""Current statevector. For num_chunks==1, the physical (chunk-sized)
|
|
222
|
+
simulator's own array. For num_chunks>1, the chunks concatenated in
|
|
223
|
+
ascending order — valid because of the MSB-first correspondence
|
|
224
|
+
between chunk index and the top `_m` logical qubits (see
|
|
225
|
+
_dispatch_multi's docstring). For the disk-overflow path, streams
|
|
226
|
+
each chunk's .npy file off disk one at a time to build the same
|
|
227
|
+
concatenation -- this DOES materialize the full (2**n,) array in
|
|
228
|
+
RAM, unlike run_chunk() itself; only use it on a size you know
|
|
229
|
+
fits, e.g. for a final readout after a run."""
|
|
230
|
+
if self._chunk_paths is not None:
|
|
231
|
+
return np.concatenate([np.load(p) for p in self._chunk_paths])
|
|
232
|
+
if self._chunk_sims is None:
|
|
233
|
+
return self._inner_sim.sv
|
|
234
|
+
xp = self._chunk_sims[0].xp
|
|
235
|
+
return xp.concatenate([sim.sv for sim in self._chunk_sims])
|
|
236
|
+
|
|
237
|
+
@sv.setter
|
|
238
|
+
def sv(self, value):
|
|
239
|
+
"""Accepts a full-length (2**n,) statevector (e.g. the output of
|
|
240
|
+
NoiseModel.apply_to_sv called on `.sv`) and writes it back through
|
|
241
|
+
to the physical storage -- the inner simulator directly for
|
|
242
|
+
num_chunks==1, split back into per-chunk slices (same ascending
|
|
243
|
+
concatenation order as the getter) for num_chunks>1, or rewritten
|
|
244
|
+
to each chunk's .npy file for the disk-overflow path."""
|
|
245
|
+
chunk_dim = self._mem_chunker.chunk_dim
|
|
246
|
+
if self._chunk_paths is not None:
|
|
247
|
+
value = np.asarray(value)
|
|
248
|
+
for i, path in enumerate(self._chunk_paths):
|
|
249
|
+
np.save(path, value[i * chunk_dim:(i + 1) * chunk_dim])
|
|
250
|
+
return
|
|
251
|
+
if self._chunk_sims is None:
|
|
252
|
+
self._inner_sim.sv = value
|
|
253
|
+
return
|
|
254
|
+
xp = self._chunk_sims[0].xp
|
|
255
|
+
value = xp.asarray(value)
|
|
256
|
+
for i, sim in enumerate(self._chunk_sims):
|
|
257
|
+
sim.sv = value[i * chunk_dim:(i + 1) * chunk_dim]
|
|
258
|
+
|
|
259
|
+
def memory_mb(self) -> float:
|
|
260
|
+
"""RAM used by the physical statevector(s) in MB -- 0 for the
|
|
261
|
+
disk-overflow path (see memory_geometry.memory_mb() for the
|
|
262
|
+
per-chunk on-disk size instead, and disk_overflow.py's own
|
|
263
|
+
docstring for why nothing (2**n_qubits,)-sized, or even
|
|
264
|
+
num_chunks-chunks-sized, is ever resident in RAM at once)."""
|
|
265
|
+
if self._chunk_paths is not None:
|
|
266
|
+
return 0.0
|
|
267
|
+
if self._chunk_sims is None:
|
|
268
|
+
return self._inner_sim.memory_mb()
|
|
269
|
+
return sum(sim.memory_mb() for sim in self._chunk_sims)
|
|
270
|
+
|
|
271
|
+
def get_probabilities(self):
|
|
272
|
+
"""|amplitude|^2 for every basis state.
|
|
273
|
+
|
|
274
|
+
num_chunks==1: forwards to the inner DenseSVSimulator for parity
|
|
275
|
+
with its own get_probabilities().
|
|
276
|
+
|
|
277
|
+
num_chunks>1: concatenates the RAW statevectors first and normalizes
|
|
278
|
+
ONCE over the full array — NOT each chunk's own get_probabilities()
|
|
279
|
+
(that would independently renormalize each chunk's partial mass to
|
|
280
|
+
1, summing to num_chunks overall and destroying the relative
|
|
281
|
+
weighting between chunks). Disk-overflow path: same normalization,
|
|
282
|
+
chunks streamed from their .npy files instead of live arrays."""
|
|
283
|
+
if self._chunk_paths is not None:
|
|
284
|
+
full_sv = np.concatenate([np.load(p) for p in self._chunk_paths])
|
|
285
|
+
elif self._chunk_sims is None:
|
|
286
|
+
return self._inner_sim.get_probabilities()
|
|
287
|
+
else:
|
|
288
|
+
full_sv = np.concatenate([np.array(sim.sv) for sim in self._chunk_sims])
|
|
289
|
+
probs = np.abs(full_sv) ** 2
|
|
290
|
+
probs = np.clip(probs, 0.0, 1.0)
|
|
291
|
+
total = probs.sum()
|
|
292
|
+
if total > 1e-12:
|
|
293
|
+
probs /= total
|
|
294
|
+
return probs
|
|
295
|
+
|
|
296
|
+
def get_statevector(self):
|
|
297
|
+
"""Full complex statevector, num_qubits logical qubits long
|
|
298
|
+
(2**n elements). num_chunks==1: forwards to the inner
|
|
299
|
+
DenseSVSimulator. num_chunks>1: raw chunks concatenated in order
|
|
300
|
+
(see `sv` property). Disk-overflow path: streamed from the .npy
|
|
301
|
+
files, same order."""
|
|
302
|
+
if self._chunk_paths is not None:
|
|
303
|
+
return np.concatenate([np.load(p) for p in self._chunk_paths])
|
|
304
|
+
if self._chunk_sims is None:
|
|
305
|
+
return self._inner_sim.get_statevector()
|
|
306
|
+
return np.concatenate([np.array(sim.sv, dtype=sim.dtype) for sim in self._chunk_sims])
|
|
307
|
+
|
|
308
|
+
# ── Multi-chunk gate dispatch (num_chunks > 1) ───────────────────
|
|
309
|
+
|
|
310
|
+
def _dispatch_multi(self, circuit: List) -> None:
|
|
311
|
+
"""Executes *circuit* against the num_chunks>1 chunk representation
|
|
312
|
+
via one jax.lax.scan call over the whole (transpiled, GATE_IDS-
|
|
313
|
+
compiled) circuit — see _build_multi_chunk_step/_build_multi_chunk_runner
|
|
314
|
+
above for the kernel, and _compile_multi_chunk_ops for the encoding.
|
|
315
|
+
|
|
316
|
+
Convention: chunk index `c` (m = self._m bits, MSB-first, same
|
|
317
|
+
n-1-qubit convention as DenseSVSimulator) equals the value of the
|
|
318
|
+
top m logical qubits (indices [0, m)); chunk_sims[c].sv holds the
|
|
319
|
+
chunk_dim amplitudes for the remaining (local) qubits [m, n). This
|
|
320
|
+
makes full_sv.reshape(num_chunks, chunk_dim)[c] == chunk_sims[c].sv
|
|
321
|
+
exactly, since NumPy's row-major reshape splits a (2,)*n tensor on
|
|
322
|
+
the leading axes first — i.e. the most-significant qubits, matching
|
|
323
|
+
this simulator's MSB-first indexing throughout. Stacking
|
|
324
|
+
chunk_sims[i].sv into one (num_chunks, chunk_dim) array before the
|
|
325
|
+
scan, and unstacking after, holds exactly the same total element
|
|
326
|
+
count as the num_chunks separate arrays it replaces — the anti-OOM
|
|
327
|
+
property this class exists for is preserved, nothing (2**n_qubits,)
|
|
328
|
+
shaped is ever materialized."""
|
|
329
|
+
compiled_ops = _compile_multi_chunk_ops(circuit)
|
|
330
|
+
xp = self._chunk_sims[0].xp
|
|
331
|
+
stacked = xp.stack([sim.sv for sim in self._chunk_sims])
|
|
332
|
+
final = self._multi_chunk_runner(stacked, compiled_ops)
|
|
333
|
+
for i, sim in enumerate(self._chunk_sims):
|
|
334
|
+
sim.sv = final[i]
|
|
335
|
+
|
|
336
|
+
# ── Distributed multi-device gate dispatch (issue #1) ────────────
|
|
337
|
+
|
|
338
|
+
def dispatch_distributed(self, circuit: List) -> None:
|
|
339
|
+
"""Executes *circuit* the same way _dispatch_multi does, but with
|
|
340
|
+
each chunk pinned to its own physical JAX device instead of all
|
|
341
|
+
chunks sharing one process's RAM — see
|
|
342
|
+
_build_distributed_chunk_step/_build_distributed_chunk_runner for
|
|
343
|
+
the kernel. Requires jax.device_count() >= num_chunks (v1 scope:
|
|
344
|
+
exactly one chunk per device); raises RuntimeError otherwise
|
|
345
|
+
rather than silently falling back to the single-process path,
|
|
346
|
+
since that would silently give up the whole point of calling this
|
|
347
|
+
method instead of run_chunk().
|
|
348
|
+
|
|
349
|
+
The (num_chunks, chunk_dim) logical array is never materialized
|
|
350
|
+
on any single device here (unlike _dispatch_multi, where it's one
|
|
351
|
+
process's RAM) -- each device holds and ever sees only its own
|
|
352
|
+
(chunk_dim,) row, exchanging edge data with its stride-partner
|
|
353
|
+
device via jax.lax.ppermute inside the kernel, not through this
|
|
354
|
+
Python method."""
|
|
355
|
+
import jax
|
|
356
|
+
if self._chunk_sims is None:
|
|
357
|
+
raise RuntimeError(
|
|
358
|
+
"dispatch_distributed() requires num_chunks > 1 "
|
|
359
|
+
"(this Chunk instance fits in a single chunk)."
|
|
360
|
+
)
|
|
361
|
+
num_chunks = self._mem_chunker.num_chunks
|
|
362
|
+
available = jax.device_count()
|
|
363
|
+
if available < num_chunks:
|
|
364
|
+
raise RuntimeError(
|
|
365
|
+
f"dispatch_distributed() needs >= {num_chunks} JAX devices "
|
|
366
|
+
f"(one per chunk), only {available} available. Force extra "
|
|
367
|
+
f"CPU devices for testing via the XLA_FLAGS environment "
|
|
368
|
+
f"variable: --xla_force_host_platform_device_count=N "
|
|
369
|
+
f"(set before the process starts, JAX's device count is "
|
|
370
|
+
f"fixed at first initialization)."
|
|
371
|
+
)
|
|
372
|
+
if self._distributed_runner is None:
|
|
373
|
+
self._distributed_runner, self._distributed_mesh = _build_distributed_chunk_runner(
|
|
374
|
+
num_chunks, self._m, self._mem_chunker.chunk_size_bits)
|
|
375
|
+
|
|
376
|
+
compiled_ops = _compile_multi_chunk_ops(circuit)
|
|
377
|
+
xp = self._chunk_sims[0].xp
|
|
378
|
+
stacked = xp.stack([sim.sv for sim in self._chunk_sims])
|
|
379
|
+
final = self._distributed_runner(stacked, compiled_ops)
|
|
380
|
+
for i, sim in enumerate(self._chunk_sims):
|
|
381
|
+
sim.sv = np.asarray(final[i])
|
|
382
|
+
|
|
383
|
+
# ── Public API ───────────────────────────────────────────
|
|
384
|
+
|
|
385
|
+
def run_chunk(
|
|
386
|
+
self,
|
|
387
|
+
circuit: List,
|
|
388
|
+
chunk_size_gates: Optional[int] = None,
|
|
389
|
+
) -> None:
|
|
390
|
+
|
|
391
|
+
if self._chunk_paths is not None:
|
|
392
|
+
compiled_ops = _compile_multi_chunk_ops(circuit)
|
|
393
|
+
run_disk_overflow_circuit(
|
|
394
|
+
self._chunk_paths, compiled_ops, self._m, self._mem_chunker.chunk_size_bits)
|
|
395
|
+
return
|
|
396
|
+
if self._chunk_sims is not None:
|
|
397
|
+
self._dispatch_multi(circuit)
|
|
398
|
+
return
|
|
399
|
+
size = chunk_size_gates if chunk_size_gates is not None else self.chunk_size_gates
|
|
400
|
+
self._circuit_chunker.split_circuit(circuit, chunk_size=size)
|
|
401
|
+
|
|
402
|
+
def run_chunk_distributed(self, circuit: List) -> None:
|
|
403
|
+
"""Like run_chunk(), but dispatches across a real JAX device mesh
|
|
404
|
+
(dispatch_distributed) instead of one process's RAM — issue #1.
|
|
405
|
+
Requires jax.device_count() >= num_chunks; raises RuntimeError
|
|
406
|
+
otherwise (see dispatch_distributed's docstring for how to test
|
|
407
|
+
this with simulated multi-device CPU)."""
|
|
408
|
+
self.dispatch_distributed(circuit)
|
|
409
|
+
|
|
410
|
+
def __repr__(self) -> str:
|
|
411
|
+
s = self._guard.status()
|
|
412
|
+
safe_qubits = self._inner_sim.n if self._inner_sim is not None else self._mem_chunker.chunk_size_bits
|
|
413
|
+
storage = f"disk ({self._disk_dir})" if self._chunk_paths is not None else "ram"
|
|
414
|
+
return (
|
|
415
|
+
f"Chunk(n_qubits={self.n}, "
|
|
416
|
+
f"safe_qubits={safe_qubits}, "
|
|
417
|
+
f"num_chunks={self.num_chunks}, "
|
|
418
|
+
f"chunk_size_bits={self.chunk_size_bits}, "
|
|
419
|
+
f"storage={storage}, "
|
|
420
|
+
f"dtype={self.dtype}, "
|
|
421
|
+
f"mem_per_chunk={self.memory_mb():.1f} MB, "
|
|
422
|
+
f"ram_free={s['free_pct']:.1f}%, "
|
|
423
|
+
f"has_jax={HAS_JAX})"
|
|
424
|
+
)
|
|
425
|
+
|
|
426
|
+
|
|
427
|
+
# ─────────────────────────────────────────────────────────────────────────────────
|
|
428
|
+
# Backward-compatibility aliases
|
|
429
|
+
# ─────────────────────────────────────────────────────────────────────────────────
|
|
430
|
+
chunk1 = MemoryChunker
|
|
431
|
+
chunk2 = Chunk
|
|
432
|
+
Chunk2Incrociato = Chunk
|
|
@@ -0,0 +1,232 @@
|
|
|
1
|
+
"""Disk-backed statevector overflow for Chunk -- Pednault et al. 2019
|
|
2
|
+
(arXiv:1910.09534, "Leveraging Secondary Storage to Simulate Deep
|
|
3
|
+
54-qubit Sycamore Circuits"): when num_chunks chunks don't all fit in
|
|
4
|
+
RAM at once, keep the idle ones on disk as plain .npy files and only
|
|
5
|
+
materialize, as a jax.Array, the small working set an individual gate
|
|
6
|
+
actually needs -- one chunk for a local gate, two (an XOR-stride pair)
|
|
7
|
+
for a chunk-mixing gate -- instead of the whole (num_chunks, chunk_dim)
|
|
8
|
+
stack _dispatch_multi holds in RAM for the entire circuit.
|
|
9
|
+
|
|
10
|
+
v1 scope (correctness-first, matching run_chunk_distributed's own staged
|
|
11
|
+
rollout): processes one chunk / one pair at a time, no batching multiple
|
|
12
|
+
pairs into one call for speed. This is strictly slower per gate than the
|
|
13
|
+
in-RAM path -- it exists to make otherwise-impossible sizes possible at
|
|
14
|
+
all, not to compete with it on speed. See docs/api/chunk.md for the real
|
|
15
|
+
distinction and a verified demo.
|
|
16
|
+
|
|
17
|
+
Every gate is one of the same 6 cases kernels.py's in-RAM
|
|
18
|
+
_build_multi_chunk_step already classifies. This module reuses that
|
|
19
|
+
classification (not a fresh derivation) via the same needs_comm split
|
|
20
|
+
_build_distributed_chunk_step's own docstring documents:
|
|
21
|
+
- 1-qubit, q1 >= m -> LOCAL: needs only this one chunk.
|
|
22
|
+
- 2-qubit, q1 >= m and q2 >= m -> LOCAL: needs only this one chunk.
|
|
23
|
+
- 2-qubit, q1 < m, q2 >= m ("ctrl chunk-select, tgt local") -> LOCAL,
|
|
24
|
+
but conditional on this chunk's OWN absolute index (no partner
|
|
25
|
+
needed -- kernels.py's _case_2q_ctrl_chunk_tgt_local already
|
|
26
|
+
computes exactly this from an explicit `idxc`).
|
|
27
|
+
- everything else (1-qubit q1 < m, or 2-qubit q2 < m) -> MIX: needs
|
|
28
|
+
exactly the XOR-stride partner chunk.
|
|
29
|
+
"""
|
|
30
|
+
import numpy as np
|
|
31
|
+
import jax.numpy as jnp
|
|
32
|
+
|
|
33
|
+
from .kernels import (
|
|
34
|
+
_gate_matrix_elements, _case_1q_local, _case_2q_local_local,
|
|
35
|
+
_case_2q_ctrl_chunk_tgt_local,
|
|
36
|
+
)
|
|
37
|
+
|
|
38
|
+
__all__ = ["partition_ops_into_phases", "run_disk_overflow_circuit", "LocalPhase", "ConditionalPhase", "MixPhase"]
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
class LocalPhase:
|
|
42
|
+
"""A maximal run of consecutive gates with q1,q2 >= m (or unused) --
|
|
43
|
+
applied to each chunk independently, no partner chunk ever needed."""
|
|
44
|
+
__slots__ = ("ops",)
|
|
45
|
+
|
|
46
|
+
def __init__(self, ops):
|
|
47
|
+
self.ops = ops # list of (g_id, q1, q2, param) python tuples
|
|
48
|
+
|
|
49
|
+
|
|
50
|
+
class ConditionalPhase:
|
|
51
|
+
"""A maximal run of consecutive 2-qubit gates with ctrl chunk-select
|
|
52
|
+
(q1 < m), tgt local (q2 >= m) -- applied to each chunk independently,
|
|
53
|
+
each op conditioned on that chunk's own absolute index (see
|
|
54
|
+
_case_2q_ctrl_chunk_tgt_local). Batched the same way LocalPhase is
|
|
55
|
+
(was: strictly one gate per phase object, so consecutive
|
|
56
|
+
ConditionalPhase-eligible gates did one full disk load/save cycle per
|
|
57
|
+
chunk PER GATE instead of once for the whole run -- prog.txt point 5b)."""
|
|
58
|
+
__slots__ = ("ops",)
|
|
59
|
+
|
|
60
|
+
def __init__(self, ops):
|
|
61
|
+
self.ops = ops # list of (g_id, q1, q2, param) python tuples
|
|
62
|
+
|
|
63
|
+
|
|
64
|
+
class MixPhase:
|
|
65
|
+
"""A single gate that needs a partner chunk: 1-qubit with q1 < m, or
|
|
66
|
+
2-qubit with q2 < m (covers both "ctrl local/tgt chunk" and "ctrl
|
|
67
|
+
AND tgt both chunk-select"). `stride` is the mixing qubit (q1 for the
|
|
68
|
+
1-qubit case, q2 for the 2-qubit cases) -- chunk i always pairs with
|
|
69
|
+
chunk i ^ (1 << (m - 1 - stride))."""
|
|
70
|
+
__slots__ = ("op", "stride")
|
|
71
|
+
|
|
72
|
+
def __init__(self, op, stride):
|
|
73
|
+
self.op = op
|
|
74
|
+
self.stride = stride
|
|
75
|
+
|
|
76
|
+
|
|
77
|
+
def partition_ops_into_phases(compiled_ops, m: int):
|
|
78
|
+
"""compiled_ops: the (n_gates, 4) [g_id, q1, q2, param] array
|
|
79
|
+
_compile_multi_chunk_ops already builds (reused as-is). Returns a
|
|
80
|
+
list of LocalPhase/ConditionalPhase/MixPhase in circuit order. Pure
|
|
81
|
+
Python -- runs once per run_chunk() call, not traced/jitted."""
|
|
82
|
+
rows = np.asarray(compiled_ops)
|
|
83
|
+
phases = []
|
|
84
|
+
pending_local = []
|
|
85
|
+
pending_conditional = []
|
|
86
|
+
|
|
87
|
+
def flush_local():
|
|
88
|
+
if pending_local:
|
|
89
|
+
phases.append(LocalPhase(list(pending_local)))
|
|
90
|
+
pending_local.clear()
|
|
91
|
+
|
|
92
|
+
def flush_conditional():
|
|
93
|
+
if pending_conditional:
|
|
94
|
+
phases.append(ConditionalPhase(list(pending_conditional)))
|
|
95
|
+
pending_conditional.clear()
|
|
96
|
+
|
|
97
|
+
for row in rows:
|
|
98
|
+
g_id, q1, q2, param = int(row[0]), int(row[1]), int(row[2]), float(row[3])
|
|
99
|
+
is_2q = g_id >= 20
|
|
100
|
+
if not is_2q:
|
|
101
|
+
if q1 < m:
|
|
102
|
+
flush_local()
|
|
103
|
+
flush_conditional()
|
|
104
|
+
phases.append(MixPhase((g_id, q1, q2, param), q1))
|
|
105
|
+
else:
|
|
106
|
+
flush_conditional()
|
|
107
|
+
pending_local.append((g_id, q1, q2, param))
|
|
108
|
+
elif q1 < m and q2 >= m:
|
|
109
|
+
flush_local()
|
|
110
|
+
pending_conditional.append((g_id, q1, q2, param))
|
|
111
|
+
elif q2 < m:
|
|
112
|
+
flush_local()
|
|
113
|
+
flush_conditional()
|
|
114
|
+
phases.append(MixPhase((g_id, q1, q2, param), q2))
|
|
115
|
+
else:
|
|
116
|
+
flush_conditional()
|
|
117
|
+
pending_local.append((g_id, q1, q2, param))
|
|
118
|
+
flush_local()
|
|
119
|
+
flush_conditional()
|
|
120
|
+
return phases
|
|
121
|
+
|
|
122
|
+
|
|
123
|
+
def _run_local_phase_on_chunk(chunk_arr, ops, m: int, k: int):
|
|
124
|
+
"""chunk_arr: (chunk_dim,) array for one chunk. Applies every op via
|
|
125
|
+
the same _case_1q_local/_case_2q_local_local formulas the in-RAM
|
|
126
|
+
kernel uses, as a (1, chunk_dim) batch of size 1 -- both cases never
|
|
127
|
+
reference the batch dimension's meaning, so this is exact, not an
|
|
128
|
+
approximation of the vectorized path."""
|
|
129
|
+
c = chunk_arr[None, :]
|
|
130
|
+
dtype = c.dtype
|
|
131
|
+
for g_id, q1, q2, param in ops:
|
|
132
|
+
g00, g01, g10, g11, u00, u01, u10, u11 = _gate_matrix_elements(g_id, param, dtype)
|
|
133
|
+
if g_id >= 20:
|
|
134
|
+
c = _case_2q_local_local(c, u00, u01, u10, u11, q1, q2, m, k)
|
|
135
|
+
else:
|
|
136
|
+
c = _case_1q_local(c, g00, g01, g10, g11, q1, m, k)
|
|
137
|
+
return c[0]
|
|
138
|
+
|
|
139
|
+
|
|
140
|
+
def _run_conditional_phase_on_chunk(chunk_arr, ops, chunk_index: int, m: int, k: int):
|
|
141
|
+
"""A run of ConditionalPhase gates, applied to one chunk whose real
|
|
142
|
+
absolute index is `chunk_index` (needed for the ctrl-bit decision;
|
|
143
|
+
see _case_2q_ctrl_chunk_tgt_local's own docstring). Same load-once/
|
|
144
|
+
apply-all/save-once shape as _run_local_phase_on_chunk."""
|
|
145
|
+
dtype = chunk_arr.dtype
|
|
146
|
+
c = chunk_arr[None, :]
|
|
147
|
+
idxc = jnp.array([chunk_index], dtype=jnp.int32)
|
|
148
|
+
for g_id, q1, q2, param in ops:
|
|
149
|
+
_, _, _, _, u00, u01, u10, u11 = _gate_matrix_elements(g_id, param, dtype)
|
|
150
|
+
c = _case_2q_ctrl_chunk_tgt_local(c, u00, u01, u10, u11, q1, q2, m, k, idxc)
|
|
151
|
+
return c[0]
|
|
152
|
+
|
|
153
|
+
|
|
154
|
+
def _mix_pair(row_a, row_b, e00, e01, e10, e11):
|
|
155
|
+
"""The 2x2 amplitude-mixing algebra shared by every chunk-select-
|
|
156
|
+
qubit case in kernels.py (_case_1q_chunk / _case_2q_ctrl_local_tgt_chunk
|
|
157
|
+
/ _case_2q_both_chunk): row_a is the row on the "mask/is_c0 = True"
|
|
158
|
+
side, row_b its XOR-stride partner. Mirrors those cases' own
|
|
159
|
+
new0/new1 formula exactly -- new0 = e00*_c + e01*amp_pair evaluated
|
|
160
|
+
at the True-side row (_c=row_a, amp_pair=row_b); new1 = e10*amp_pair
|
|
161
|
+
+ e11*_c evaluated at the False-side row, where _c=row_b and
|
|
162
|
+
amp_pair=row_a, i.e. new_b = e10*row_a + e11*row_b."""
|
|
163
|
+
new_a = e00 * row_a + e01 * row_b
|
|
164
|
+
new_b = e10 * row_a + e11 * row_b
|
|
165
|
+
return new_a, new_b
|
|
166
|
+
|
|
167
|
+
|
|
168
|
+
def _run_mix_phase_on_pair(row_a, row_b, op, index_a: int, m: int, k: int):
|
|
169
|
+
"""row_a is the chunk whose bit at the gate's own mixing qubit is 0,
|
|
170
|
+
row_b its XOR-stride partner (bit 1) -- the caller picks which
|
|
171
|
+
loaded chunk is "a" vs "b" by checking that bit on the real absolute
|
|
172
|
+
indices before calling this. `index_a` is row_a's real absolute
|
|
173
|
+
chunk index (needed only for case 6's ctrl-chunk-select decision)."""
|
|
174
|
+
g_id, q1, q2, param = op
|
|
175
|
+
dtype = row_a.dtype
|
|
176
|
+
g00, g01, g10, g11, u00, u01, u10, u11 = _gate_matrix_elements(g_id, param, dtype)
|
|
177
|
+
if g_id < 20:
|
|
178
|
+
# case 2: 1-qubit, chunk-select q1 -- unconditional mix.
|
|
179
|
+
return _mix_pair(row_a, row_b, g00, g01, g10, g11)
|
|
180
|
+
if q1 >= m:
|
|
181
|
+
# case 5: ctrl LOCAL (elementwise mask within the row), tgt
|
|
182
|
+
# chunk-select q2 (the mixing qubit).
|
|
183
|
+
ctrl_phys = (k - 1) - (q1 - m)
|
|
184
|
+
idxl = jnp.arange(1 << k, dtype=jnp.int32)
|
|
185
|
+
ctrl_bit = (idxl & (jnp.int32(1) << ctrl_phys)) != 0
|
|
186
|
+
new_a, new_b = _mix_pair(row_a, row_b, u00, u01, u10, u11)
|
|
187
|
+
return jnp.where(ctrl_bit, new_a, row_a), jnp.where(ctrl_bit, new_b, row_b)
|
|
188
|
+
# case 6: ctrl AND tgt both chunk-select; mixing qubit is q2, ctrl
|
|
189
|
+
# decided once from index_a's own bit at q1 (identical for both rows
|
|
190
|
+
# of the pair -- XOR-ing the tgt/q2 stride never touches the q1 bit).
|
|
191
|
+
ctrl_stride = 1 << (m - 1 - q1)
|
|
192
|
+
ctrl_set = (index_a & ctrl_stride) != 0
|
|
193
|
+
new_a, new_b = _mix_pair(row_a, row_b, u00, u01, u10, u11)
|
|
194
|
+
return (new_a, new_b) if ctrl_set else (row_a, row_b)
|
|
195
|
+
|
|
196
|
+
|
|
197
|
+
def run_disk_overflow_circuit(chunk_paths, compiled_ops, m: int, k: int):
|
|
198
|
+
"""Runs the compiled circuit against num_chunks chunks stored as
|
|
199
|
+
plain .npy files at `chunk_paths` (index = real chunk index),
|
|
200
|
+
mutating those files in place. Never materializes more than one
|
|
201
|
+
(LocalPhase/ConditionalPhase) or two (MixPhase) chunks as jax.Array
|
|
202
|
+
at once, regardless of num_chunks -- this is the whole point."""
|
|
203
|
+
num_chunks = len(chunk_paths)
|
|
204
|
+
phases = partition_ops_into_phases(compiled_ops, m)
|
|
205
|
+
|
|
206
|
+
for phase in phases:
|
|
207
|
+
if isinstance(phase, LocalPhase):
|
|
208
|
+
for i, path in enumerate(chunk_paths):
|
|
209
|
+
arr = jnp.asarray(np.load(path))
|
|
210
|
+
arr = _run_local_phase_on_chunk(arr, phase.ops, m, k)
|
|
211
|
+
np.save(path, np.asarray(arr))
|
|
212
|
+
|
|
213
|
+
elif isinstance(phase, ConditionalPhase):
|
|
214
|
+
for i, path in enumerate(chunk_paths):
|
|
215
|
+
arr = jnp.asarray(np.load(path))
|
|
216
|
+
arr = _run_conditional_phase_on_chunk(arr, phase.ops, i, m, k)
|
|
217
|
+
np.save(path, np.asarray(arr))
|
|
218
|
+
|
|
219
|
+
elif isinstance(phase, MixPhase):
|
|
220
|
+
stride_bits = 1 << (m - 1 - phase.stride)
|
|
221
|
+
done = set()
|
|
222
|
+
for i in range(num_chunks):
|
|
223
|
+
if i in done:
|
|
224
|
+
continue
|
|
225
|
+
j = i ^ stride_bits
|
|
226
|
+
done.add(i)
|
|
227
|
+
done.add(j)
|
|
228
|
+
row_a = jnp.asarray(np.load(chunk_paths[i]))
|
|
229
|
+
row_b = jnp.asarray(np.load(chunk_paths[j]))
|
|
230
|
+
new_a, new_b = _run_mix_phase_on_pair(row_a, row_b, phase.op, i, m, k)
|
|
231
|
+
np.save(chunk_paths[i], np.asarray(new_a))
|
|
232
|
+
np.save(chunk_paths[j], np.asarray(new_b))
|