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,531 @@
|
|
|
1
|
+
import functools
|
|
2
|
+
from typing import List
|
|
3
|
+
|
|
4
|
+
import jax
|
|
5
|
+
import jax.numpy as jnp
|
|
6
|
+
|
|
7
|
+
from ._engine_imports import QuantumTranspiler, GATE_IDS, _gate_1q_matrix
|
|
8
|
+
|
|
9
|
+
__all__ = [
|
|
10
|
+
"_build_multi_chunk_step", "_build_multi_chunk_runner",
|
|
11
|
+
"_build_distributed_chunk_step", "_build_distributed_chunk_runner",
|
|
12
|
+
"_compile_multi_chunk_ops",
|
|
13
|
+
# Re-exported for disk_overflow.py's phased/streaming path -- see its
|
|
14
|
+
# module docstring for which of these it calls and why.
|
|
15
|
+
"_gate_matrix_elements", "_case_1q_local", "_case_2q_local_local",
|
|
16
|
+
"_case_2q_ctrl_chunk_tgt_local",
|
|
17
|
+
]
|
|
18
|
+
|
|
19
|
+
# ─────────────────────────────────────────────────────────────────────────────────
|
|
20
|
+
# Multi-chunk JIT kernel (num_chunks > 1)
|
|
21
|
+
# ─────────────────────────────────────────────────────────────────────────────────
|
|
22
|
+
|
|
23
|
+
# ─────────────────────────────────────────────────────────────────────────────────
|
|
24
|
+
# The 6 gate/qubit-location case bodies, factored out to module level so
|
|
25
|
+
# disk_overflow.py's phased/streaming path can call the exact same verified
|
|
26
|
+
# formulas instead of re-deriving them -- see that module's docstring for
|
|
27
|
+
# which of these 6 need a partner chunk (MixPhase) vs. run standalone per
|
|
28
|
+
# chunk (LocalPhase), a classification already worked out once for
|
|
29
|
+
# _build_distributed_chunk_step below (needs_comm_1q/needs_comm_2q) and
|
|
30
|
+
# reused rather than re-derived a third time.
|
|
31
|
+
#
|
|
32
|
+
# Verified case-by-case against DenseSVSimulator on non-chunked reference
|
|
33
|
+
# circuits before being wired in, not derived fresh -- this extraction is
|
|
34
|
+
# a pure move, not a rewrite: same array shapes, same operations, same
|
|
35
|
+
# order, only the enclosing scope changed from a closure to explicit args.
|
|
36
|
+
# ─────────────────────────────────────────────────────────────────────────────────
|
|
37
|
+
|
|
38
|
+
def _gate_matrix_elements(g_id, param, dtype):
|
|
39
|
+
"""1-qubit matrix elements (g00,g01,g10,g11) and 2-qubit controlled-U
|
|
40
|
+
submatrix elements (u00,u01,u10,u11) for one compiled [g_id, param]
|
|
41
|
+
pair. Factored out of _build_multi_chunk_step's step() so
|
|
42
|
+
disk_overflow.py's per-gate phase processors compute matrices from
|
|
43
|
+
this SAME table instead of a third hand-copied one — g_id/param may
|
|
44
|
+
be traced (the in-RAM scan) or concrete Python/jnp scalars (a single
|
|
45
|
+
known gate in a phase); jax.lax.switch works identically either way.
|
|
46
|
+
|
|
47
|
+
The 1-qubit table itself lives in compiler.py's _gate_1q_matrix now
|
|
48
|
+
(was: an independent hand-copied switch here, with a "must stay in
|
|
49
|
+
sync" comment -- prog.txt point 2); only the 2-qubit controlled-U
|
|
50
|
+
table below is local to this module."""
|
|
51
|
+
half_p = param * jnp.float64(0.5)
|
|
52
|
+
exp_pos = jnp.exp(1j * param).astype(dtype)
|
|
53
|
+
exp_pos_half = jnp.exp(1j * half_p).astype(dtype)
|
|
54
|
+
exp_neg_half = jnp.exp(-1j * half_p).astype(dtype)
|
|
55
|
+
|
|
56
|
+
g_id = jnp.asarray(g_id).astype(jnp.int32)
|
|
57
|
+
g_1q = _gate_1q_matrix(g_id, param, dtype)
|
|
58
|
+
|
|
59
|
+
# Controlled-U submatrix for the 5 two-qubit gate types (mat[2:,2:]
|
|
60
|
+
# of each gate's full 4x4 form — same values _apply_gate_multi's
|
|
61
|
+
# `U = mat[2:, 2:]` extracted from GATES/PARAMETRIC_GATES).
|
|
62
|
+
# 20=CX->X, 21=CZ->Z, 22=CP->P(theta), 24=CY->Y, 25=CRZ->RZ(theta).
|
|
63
|
+
two_q_idx = jnp.where(g_id == 20, 0,
|
|
64
|
+
jnp.where(g_id == 21, 1,
|
|
65
|
+
jnp.where(g_id == 22, 2,
|
|
66
|
+
jnp.where(g_id == 24, 3, 4))))
|
|
67
|
+
U = jax.lax.switch(
|
|
68
|
+
two_q_idx,
|
|
69
|
+
[
|
|
70
|
+
lambda _: jnp.array([[0.0 + 0j, 1.0 + 0j], [1.0 + 0j, 0.0 + 0j]], dtype=dtype),
|
|
71
|
+
lambda _: jnp.array([[1.0 + 0j, 0.0 + 0j], [0.0 + 0j, -1.0 + 0j]], dtype=dtype),
|
|
72
|
+
lambda _: jnp.array([[1.0 + 0j, 0.0 + 0j], [0.0 + 0j, exp_pos]], dtype=dtype),
|
|
73
|
+
lambda _: jnp.array([[0.0 + 0j, -1j], [1j, 0.0 + 0j]], dtype=dtype),
|
|
74
|
+
lambda _: jnp.array([[exp_neg_half, 0.0 + 0j], [0.0 + 0j, exp_pos_half]], dtype=dtype),
|
|
75
|
+
],
|
|
76
|
+
operand=None,
|
|
77
|
+
)
|
|
78
|
+
return g_1q[0, 0], g_1q[0, 1], g_1q[1, 0], g_1q[1, 1], U[0, 0], U[0, 1], U[1, 0], U[1, 1]
|
|
79
|
+
|
|
80
|
+
|
|
81
|
+
def _case_1q_local(_c, g00, g01, g10, g11, q1, m, k):
|
|
82
|
+
"""1-qubit gate, LOCAL qubit (q1 >= m). No cross-chunk data needed --
|
|
83
|
+
valid for any batch size along axis 0, including a single chunk."""
|
|
84
|
+
local_phys = (k - 1) - (q1 - m)
|
|
85
|
+
stride = jnp.int32(1) << local_phys
|
|
86
|
+
idx = jnp.arange(1 << k, dtype=jnp.int32)
|
|
87
|
+
idx_pair = idx ^ stride
|
|
88
|
+
mask0 = (idx & stride) == 0
|
|
89
|
+
amp_pair = _c[:, idx_pair]
|
|
90
|
+
new0 = g00 * _c + g01 * amp_pair
|
|
91
|
+
new1 = g10 * amp_pair + g11 * _c
|
|
92
|
+
return jnp.where(mask0[None, :], new0, new1)
|
|
93
|
+
|
|
94
|
+
|
|
95
|
+
def _case_1q_chunk(_c, g00, g01, g10, g11, q1, m, num_chunks):
|
|
96
|
+
"""1-qubit gate, CHUNK-SELECT qubit (q1 < m). Mixes whole rows
|
|
97
|
+
pairwise across axis 0 -- needs every chunk present at once; the
|
|
98
|
+
disk-overflow path does NOT call this directly (see _mix_pair for
|
|
99
|
+
the pair-at-a-time equivalent), only the in-RAM stacked kernel does."""
|
|
100
|
+
stride = jnp.int32(1) << (m - 1 - q1)
|
|
101
|
+
idxc = jnp.arange(num_chunks, dtype=jnp.int32)
|
|
102
|
+
idxc_pair = idxc ^ stride
|
|
103
|
+
mask0 = (idxc & stride) == 0
|
|
104
|
+
amp_pair = _c[idxc_pair]
|
|
105
|
+
new0 = g00 * _c + g01 * amp_pair
|
|
106
|
+
new1 = g10 * amp_pair + g11 * _c
|
|
107
|
+
return jnp.where(mask0[:, None], new0, new1)
|
|
108
|
+
|
|
109
|
+
|
|
110
|
+
def _case_2q_local_local(_c, u00, u01, u10, u11, q1, q2, m, k):
|
|
111
|
+
"""2-qubit gate, ctrl AND tgt both LOCAL. No cross-chunk data needed
|
|
112
|
+
-- valid for any batch size along axis 0, including a single chunk."""
|
|
113
|
+
ctrl_phys = (k - 1) - (q1 - m)
|
|
114
|
+
tgt_phys = (k - 1) - (q2 - m)
|
|
115
|
+
idx = jnp.arange(1 << k, dtype=jnp.int32)
|
|
116
|
+
ctrl_bit = (idx & (jnp.int32(1) << ctrl_phys)) != 0
|
|
117
|
+
tgt_bit = (idx & (jnp.int32(1) << tgt_phys)) != 0
|
|
118
|
+
partner = idx ^ (jnp.int32(1) << tgt_phys)
|
|
119
|
+
amp_partner = _c[:, partner]
|
|
120
|
+
new0 = u00 * _c + u01 * amp_partner
|
|
121
|
+
new1 = u10 * amp_partner + u11 * _c
|
|
122
|
+
after = jnp.where(tgt_bit[None, :], new1, new0)
|
|
123
|
+
return jnp.where(ctrl_bit[None, :], after, _c)
|
|
124
|
+
|
|
125
|
+
|
|
126
|
+
def _case_2q_ctrl_chunk_tgt_local(_c, u00, u01, u10, u11, q1, q2, m, k, idxc):
|
|
127
|
+
"""ctrl CHUNK-SELECT, tgt LOCAL. Whole chunks where the chunk-index's
|
|
128
|
+
ctrl bit is set get U applied as a local 1-qubit gate; the rest are
|
|
129
|
+
untouched -- a decision each chunk makes from its OWN index alone, no
|
|
130
|
+
partner needed (same insight _build_distributed_chunk_step's docstring
|
|
131
|
+
already documents for this exact case: "no communication needed").
|
|
132
|
+
|
|
133
|
+
`idxc` is the REAL absolute chunk index of each row in `_c` (not
|
|
134
|
+
necessarily a contiguous 0..N-1 range) -- the in-RAM kernel below
|
|
135
|
+
always passes jnp.arange(num_chunks) (row i really is chunk i), but
|
|
136
|
+
disk_overflow.py's LocalPhase reuses this same function one real
|
|
137
|
+
chunk at a time, passing that chunk's own true index so this still
|
|
138
|
+
decides the ctrl bit correctly without needing any other chunk
|
|
139
|
+
resident in RAM."""
|
|
140
|
+
ctrl_stride = jnp.int32(1) << (m - 1 - q1)
|
|
141
|
+
ctrl_set = (idxc & ctrl_stride) != 0
|
|
142
|
+
tgt_phys = (k - 1) - (q2 - m)
|
|
143
|
+
idxl = jnp.arange(1 << k, dtype=jnp.int32)
|
|
144
|
+
tgt_bit = (idxl & (jnp.int32(1) << tgt_phys)) != 0
|
|
145
|
+
partner = idxl ^ (jnp.int32(1) << tgt_phys)
|
|
146
|
+
amp_partner = _c[:, partner]
|
|
147
|
+
new0 = u00 * _c + u01 * amp_partner
|
|
148
|
+
new1 = u10 * amp_partner + u11 * _c
|
|
149
|
+
after = jnp.where(tgt_bit[None, :], new1, new0)
|
|
150
|
+
return jnp.where(ctrl_set[:, None], after, _c)
|
|
151
|
+
|
|
152
|
+
|
|
153
|
+
def _case_2q_ctrl_local_tgt_chunk(_c, u00, u01, u10, u11, q1, q2, m, k, num_chunks):
|
|
154
|
+
"""ctrl LOCAL, tgt CHUNK-SELECT. Pairs of chunks get mixed, but ONLY
|
|
155
|
+
where the local ctrl bit (same position in every chunk) is set -- an
|
|
156
|
+
elementwise mask. Needs every chunk present at once; see _mix_pair
|
|
157
|
+
for the pair-at-a-time equivalent used by disk_overflow.py."""
|
|
158
|
+
ctrl_phys = (k - 1) - (q1 - m)
|
|
159
|
+
idxl = jnp.arange(1 << k, dtype=jnp.int32)
|
|
160
|
+
ctrl_bit = (idxl & (jnp.int32(1) << ctrl_phys)) != 0
|
|
161
|
+
tgt_stride = jnp.int32(1) << (m - 1 - q2)
|
|
162
|
+
idxc = jnp.arange(num_chunks, dtype=jnp.int32)
|
|
163
|
+
idxc_pair = idxc ^ tgt_stride
|
|
164
|
+
is_c0 = (idxc & tgt_stride) == 0
|
|
165
|
+
amp_pair = _c[idxc_pair]
|
|
166
|
+
new_c0 = u00 * _c + u01 * amp_pair
|
|
167
|
+
new_c1 = u10 * amp_pair + u11 * _c
|
|
168
|
+
after = jnp.where(is_c0[:, None], new_c0, new_c1)
|
|
169
|
+
return jnp.where(ctrl_bit[None, :], after, _c)
|
|
170
|
+
|
|
171
|
+
|
|
172
|
+
def _case_2q_both_chunk(_c, u00, u01, u10, u11, q1, q2, m, num_chunks):
|
|
173
|
+
"""ctrl AND tgt both CHUNK-SELECT. Needs every chunk present at once;
|
|
174
|
+
see _mix_pair for the pair-at-a-time equivalent used by disk_overflow.py."""
|
|
175
|
+
ctrl_stride = jnp.int32(1) << (m - 1 - q1)
|
|
176
|
+
tgt_stride = jnp.int32(1) << (m - 1 - q2)
|
|
177
|
+
idxc = jnp.arange(num_chunks, dtype=jnp.int32)
|
|
178
|
+
ctrl_set = (idxc & ctrl_stride) != 0
|
|
179
|
+
idxc_pair = idxc ^ tgt_stride
|
|
180
|
+
is_c0 = (idxc & tgt_stride) == 0
|
|
181
|
+
amp_pair = _c[idxc_pair]
|
|
182
|
+
new_c0 = u00 * _c + u01 * amp_pair
|
|
183
|
+
new_c1 = u10 * amp_pair + u11 * _c
|
|
184
|
+
after = jnp.where(is_c0[:, None], new_c0, new_c1)
|
|
185
|
+
return jnp.where(ctrl_set[:, None], after, _c)
|
|
186
|
+
|
|
187
|
+
|
|
188
|
+
def _build_multi_chunk_step(num_chunks: int, m: int, k: int):
|
|
189
|
+
"""
|
|
190
|
+
Build a jax.lax.scan-compatible step function operating on the
|
|
191
|
+
STACKED multi-chunk representation — a (num_chunks, chunk_dim)
|
|
192
|
+
array — instead of a flat (2**n_qubits,) statevector.
|
|
193
|
+
|
|
194
|
+
Modeled directly on compiler.py's _apply_gate_fast_step (same
|
|
195
|
+
[g_id, q1, q2, param] encoding via GATE_IDS, same lax.switch for the
|
|
196
|
+
1-qubit matrix, same per-gate controlled-U dispatch for the 5
|
|
197
|
+
two-qubit gate types) — but a gate touching a "chunk-select" qubit
|
|
198
|
+
(index < m, the top m logical qubits that select WHICH chunk) mixes
|
|
199
|
+
whole (chunk_dim,)-shaped ROWS instead of individual amplitudes,
|
|
200
|
+
while a gate touching a "local" qubit (index >= m) mixes individual
|
|
201
|
+
elements WITHIN each row in parallel. This never materializes a
|
|
202
|
+
(2**n_qubits,) array — Chunk's whole reason to exist — because the
|
|
203
|
+
stacked shape (num_chunks, chunk_dim) holds exactly the same total
|
|
204
|
+
elements as the num_chunks separate per-chunk arrays it replaces.
|
|
205
|
+
|
|
206
|
+
The 6 gate/qubit-location combinations (_case_1q_local etc. above)
|
|
207
|
+
are a direct translation of the pre-JIT _apply_gate_multi's
|
|
208
|
+
Python-loop formulas (removed once this replaced it) — verified
|
|
209
|
+
case-by-case against DenseSVSimulator on non-chunked reference
|
|
210
|
+
circuits before being wired in, not derived fresh. All 6 are traced
|
|
211
|
+
unconditionally every step and selected via jnp.where on the runtime
|
|
212
|
+
q1/q2 vs static m comparison, same "trace every branch" pattern
|
|
213
|
+
_apply_gate_fast_step already uses for is_1q/is_2q and the 5
|
|
214
|
+
two-qubit sub-gates.
|
|
215
|
+
"""
|
|
216
|
+
|
|
217
|
+
def step(chunks, operation):
|
|
218
|
+
g_id = operation[0].astype(jnp.int32)
|
|
219
|
+
q1 = operation[1].astype(jnp.int32)
|
|
220
|
+
q2 = operation[2].astype(jnp.int32)
|
|
221
|
+
param = operation[3]
|
|
222
|
+
dtype = chunks.dtype # never hardcode complex128 — see the
|
|
223
|
+
# use_float32 bug this exact mistake
|
|
224
|
+
# caused once already in beast-mode.
|
|
225
|
+
|
|
226
|
+
g00, g01, g10, g11, u00, u01, u10, u11 = _gate_matrix_elements(g_id, param, dtype)
|
|
227
|
+
|
|
228
|
+
is_2q = g_id >= 20
|
|
229
|
+
q1_chunk = q1 < m
|
|
230
|
+
q2_chunk = q2 < m
|
|
231
|
+
|
|
232
|
+
result_2q = jnp.where(
|
|
233
|
+
q1_chunk & q2_chunk, _case_2q_both_chunk(chunks, u00, u01, u10, u11, q1, q2, m, num_chunks),
|
|
234
|
+
jnp.where(q1_chunk & (~q2_chunk), _case_2q_ctrl_chunk_tgt_local(chunks, u00, u01, u10, u11, q1, q2, m, k, jnp.arange(num_chunks, dtype=jnp.int32)),
|
|
235
|
+
jnp.where((~q1_chunk) & q2_chunk, _case_2q_ctrl_local_tgt_chunk(chunks, u00, u01, u10, u11, q1, q2, m, k, num_chunks),
|
|
236
|
+
_case_2q_local_local(chunks, u00, u01, u10, u11, q1, q2, m, k))))
|
|
237
|
+
result_1q = jnp.where(
|
|
238
|
+
q1_chunk, _case_1q_chunk(chunks, g00, g01, g10, g11, q1, m, num_chunks),
|
|
239
|
+
_case_1q_local(chunks, g00, g01, g10, g11, q1, m, k))
|
|
240
|
+
|
|
241
|
+
new_chunks = jnp.where(is_2q, result_2q, result_1q)
|
|
242
|
+
return new_chunks.astype(dtype), None
|
|
243
|
+
|
|
244
|
+
return step
|
|
245
|
+
|
|
246
|
+
|
|
247
|
+
# ─────────────────────────────────────────────────────────
|
|
248
|
+
# Distributed (multi-device) variant — one physical chunk per device
|
|
249
|
+
# ─────────────────────────────────────────────────────────
|
|
250
|
+
|
|
251
|
+
def _build_distributed_chunk_step(num_chunks: int, m: int, k: int, axis_name: str):
|
|
252
|
+
"""Same 6-case formula set as _build_multi_chunk_step, but each
|
|
253
|
+
device holds exactly ONE chunk row (chunk_dim,) instead of the
|
|
254
|
+
whole (num_chunks, chunk_dim) stack living on one device/process —
|
|
255
|
+
issue #1: distribute chunks across a device mesh (multi-GPU/
|
|
256
|
+
multi-host), not just multi-chunk within one process's RAM.
|
|
257
|
+
|
|
258
|
+
The stacked-array formulation's "mix pairs of rows across axis 0"
|
|
259
|
+
(cases 2/5/6, touching a chunk-select qubit) becomes real
|
|
260
|
+
point-to-point network communication here: jax.lax.ppermute, keyed
|
|
261
|
+
on the fixed XOR-stride pairing between chunk indices — the
|
|
262
|
+
textbook pairwise-exchange communication pattern used by
|
|
263
|
+
distributed statevector simulators (each device sends its local
|
|
264
|
+
row to its stride-partner device and receives the partner's row
|
|
265
|
+
back). Cases 1/3/4 need NO communication at all: case 1/3 are
|
|
266
|
+
purely local (both qubits live inside this device's own chunk_dim
|
|
267
|
+
index space), and case 4 (ctrl chunk-select, tgt local) is a
|
|
268
|
+
decision every device can make on its OWN chunk index alone
|
|
269
|
+
(whether ITS id has the ctrl bit set) — no data from any other
|
|
270
|
+
device is needed to decide or to apply the local tgt gate.
|
|
271
|
+
|
|
272
|
+
ppermute's `perm` argument is a communication topology and must be
|
|
273
|
+
STATIC (known at trace time) — it cannot be built from q1/q2,
|
|
274
|
+
which are traced values read from the scanned circuit array. Since
|
|
275
|
+
q1/q2 only ever range over the m possible chunk-select qubit
|
|
276
|
+
indices [0, m), every possible stride is enumerated as its own
|
|
277
|
+
statically-built ppermute call, and jax.lax.switch (traced
|
|
278
|
+
unconditionally, same "trace every branch" pattern used
|
|
279
|
+
throughout this codebase) picks the right one at runtime — plus
|
|
280
|
+
one extra identity branch for "no chunk-select qubit involved,
|
|
281
|
+
no communication needed" (verified below to be exactly cases
|
|
282
|
+
1/3/4, never 2/5/6)."""
|
|
283
|
+
chunk_dim = 1 << k
|
|
284
|
+
|
|
285
|
+
def step(local_row, operation):
|
|
286
|
+
g_id = operation[0].astype(jnp.int32)
|
|
287
|
+
q1 = operation[1].astype(jnp.int32)
|
|
288
|
+
q2 = operation[2].astype(jnp.int32)
|
|
289
|
+
param = operation[3]
|
|
290
|
+
dtype = local_row.dtype
|
|
291
|
+
|
|
292
|
+
my_id = jax.lax.axis_index(axis_name).astype(jnp.int32)
|
|
293
|
+
|
|
294
|
+
# Same table _build_multi_chunk_step uses (was: an independent
|
|
295
|
+
# hand-copied switch here -- prog.txt point 2, "must stay in
|
|
296
|
+
# sync" comment on both sides).
|
|
297
|
+
g00, g01, g10, g11, u00, u01, u10, u11 = _gate_matrix_elements(g_id, param, dtype)
|
|
298
|
+
|
|
299
|
+
is_2q = g_id >= 20
|
|
300
|
+
q1_chunk = q1 < m
|
|
301
|
+
q2_chunk = q2 < m
|
|
302
|
+
|
|
303
|
+
# ── single point-to-point exchange, done once per step ──────
|
|
304
|
+
# comm_qubit selects which stride to ppermute on: q2 (tgt) for
|
|
305
|
+
# any 2-qubit gate with a chunk-select target (cases 5 and 6),
|
|
306
|
+
# q1 for a 1-qubit gate on a chunk-select qubit (case 2),
|
|
307
|
+
# sentinel `m` (-> identity, no network traffic) for cases
|
|
308
|
+
# 1/3/4, which never need another device's data.
|
|
309
|
+
needs_comm_2q = is_2q & q2_chunk
|
|
310
|
+
needs_comm_1q = (~is_2q) & q1_chunk
|
|
311
|
+
comm_qubit = jnp.where(needs_comm_2q, q2, jnp.where(needs_comm_1q, q1, m))
|
|
312
|
+
safe_comm_idx = jnp.clip(comm_qubit, 0, m)
|
|
313
|
+
|
|
314
|
+
# NOTE: `_perm` MUST be bound as a default-argument value here
|
|
315
|
+
# (evaluated eagerly, once, at lambda-creation time inside this
|
|
316
|
+
# loop) rather than referencing `q`/`m` freely inside the
|
|
317
|
+
# lambda body -- a free reference would be looked up at CALL
|
|
318
|
+
# time via Python's normal late-binding closure semantics, by
|
|
319
|
+
# which point the loop variable `q` has already reached its
|
|
320
|
+
# final value (m-1) for every branch, silently making every
|
|
321
|
+
# ppermute use the LAST qubit's stride regardless of which
|
|
322
|
+
# branch was actually selected. Caught by exactly that
|
|
323
|
+
# symptom: only the branch for q == m-1 gave correct results.
|
|
324
|
+
ppermute_branches = [
|
|
325
|
+
(lambda _row, _perm=[(i, i ^ (1 << (m - 1 - q))) for i in range(num_chunks)]:
|
|
326
|
+
jax.lax.ppermute(_row, axis_name, perm=_perm))
|
|
327
|
+
for q in range(m)
|
|
328
|
+
] + [lambda _row: _row] # identity: no chunk-select qubit involved
|
|
329
|
+
paired_row = jax.lax.switch(safe_comm_idx, ppermute_branches, local_row)
|
|
330
|
+
|
|
331
|
+
# ── case 1: 1-qubit, LOCAL qubit (q1 >= m) — no comm ─────────
|
|
332
|
+
def case_1q_local(_row):
|
|
333
|
+
local_phys = (k - 1) - (q1 - m)
|
|
334
|
+
stride = jnp.int32(1) << local_phys
|
|
335
|
+
idx = jnp.arange(chunk_dim, dtype=jnp.int32)
|
|
336
|
+
idx_pair = idx ^ stride
|
|
337
|
+
mask0 = (idx & stride) == 0
|
|
338
|
+
amp_pair = _row[idx_pair]
|
|
339
|
+
new0 = g00 * _row + g01 * amp_pair
|
|
340
|
+
new1 = g10 * amp_pair + g11 * _row
|
|
341
|
+
return jnp.where(mask0, new0, new1)
|
|
342
|
+
|
|
343
|
+
# ── case 2: 1-qubit, CHUNK-SELECT qubit (q1 < m) ───────────
|
|
344
|
+
# paired_row already fetched via ppermute above (comm_qubit=q1).
|
|
345
|
+
def case_1q_chunk(_row):
|
|
346
|
+
stride = jnp.int32(1) << (m - 1 - q1)
|
|
347
|
+
mask0 = (my_id & stride) == 0
|
|
348
|
+
new0 = g00 * _row + g01 * paired_row
|
|
349
|
+
new1 = g10 * paired_row + g11 * _row
|
|
350
|
+
return jnp.where(mask0, new0, new1)
|
|
351
|
+
|
|
352
|
+
# ── case 3: 2-qubit, ctrl AND tgt both LOCAL — no comm ───────
|
|
353
|
+
def case_2q_local_local(_row):
|
|
354
|
+
ctrl_phys = (k - 1) - (q1 - m)
|
|
355
|
+
tgt_phys = (k - 1) - (q2 - m)
|
|
356
|
+
idx = jnp.arange(chunk_dim, dtype=jnp.int32)
|
|
357
|
+
ctrl_bit = (idx & (jnp.int32(1) << ctrl_phys)) != 0
|
|
358
|
+
partner = idx ^ (jnp.int32(1) << tgt_phys)
|
|
359
|
+
amp_partner = _row[partner]
|
|
360
|
+
new0 = u00 * _row + u01 * amp_partner
|
|
361
|
+
new1 = u10 * amp_partner + u11 * _row
|
|
362
|
+
tgt_bit = (idx & (jnp.int32(1) << tgt_phys)) != 0
|
|
363
|
+
after = jnp.where(tgt_bit, new1, new0)
|
|
364
|
+
return jnp.where(ctrl_bit, after, _row)
|
|
365
|
+
|
|
366
|
+
# ── case 4: ctrl CHUNK-SELECT, tgt LOCAL — no comm needed: ─
|
|
367
|
+
# every device decides purely from its OWN chunk index (my_id)
|
|
368
|
+
# whether the control bit is set, and if so applies U locally.
|
|
369
|
+
def case_2q_ctrl_chunk_tgt_local(_row):
|
|
370
|
+
ctrl_stride = jnp.int32(1) << (m - 1 - q1)
|
|
371
|
+
ctrl_set = (my_id & ctrl_stride) != 0
|
|
372
|
+
tgt_phys = (k - 1) - (q2 - m)
|
|
373
|
+
idx = jnp.arange(chunk_dim, dtype=jnp.int32)
|
|
374
|
+
tgt_bit = (idx & (jnp.int32(1) << tgt_phys)) != 0
|
|
375
|
+
partner = idx ^ (jnp.int32(1) << tgt_phys)
|
|
376
|
+
amp_partner = _row[partner]
|
|
377
|
+
new0 = u00 * _row + u01 * amp_partner
|
|
378
|
+
new1 = u10 * amp_partner + u11 * _row
|
|
379
|
+
after = jnp.where(tgt_bit, new1, new0)
|
|
380
|
+
return jnp.where(ctrl_set, after, _row)
|
|
381
|
+
|
|
382
|
+
# ── case 5: ctrl LOCAL, tgt CHUNK-SELECT ──────────────
|
|
383
|
+
# paired_row already fetched via ppermute above (comm_qubit=q2,
|
|
384
|
+
# keyed on tgt's stride). is_c0 is a per-device scalar decision
|
|
385
|
+
# (which side of the tgt pairing this device is on); ctrl_bit
|
|
386
|
+
# is a per-element mask within the row (elementwise, local).
|
|
387
|
+
def case_2q_ctrl_local_tgt_chunk(_row):
|
|
388
|
+
ctrl_phys = (k - 1) - (q1 - m)
|
|
389
|
+
idx = jnp.arange(chunk_dim, dtype=jnp.int32)
|
|
390
|
+
ctrl_bit = (idx & (jnp.int32(1) << ctrl_phys)) != 0
|
|
391
|
+
tgt_stride = jnp.int32(1) << (m - 1 - q2)
|
|
392
|
+
is_c0 = (my_id & tgt_stride) == 0
|
|
393
|
+
new_c0 = u00 * _row + u01 * paired_row
|
|
394
|
+
new_c1 = u10 * paired_row + u11 * _row
|
|
395
|
+
after = jnp.where(is_c0, new_c0, new_c1)
|
|
396
|
+
return jnp.where(ctrl_bit, after, _row)
|
|
397
|
+
|
|
398
|
+
# ── case 6: ctrl AND tgt both CHUNK-SELECT ────────────
|
|
399
|
+
# paired_row via ppermute keyed on tgt's stride (comm_qubit=q2);
|
|
400
|
+
# ctrl_set and is_c0 are both per-device scalars (this device's
|
|
401
|
+
# own chunk index bits) — no per-element masking needed at all.
|
|
402
|
+
def case_2q_both_chunk(_row):
|
|
403
|
+
ctrl_stride = jnp.int32(1) << (m - 1 - q1)
|
|
404
|
+
ctrl_set = (my_id & ctrl_stride) != 0
|
|
405
|
+
tgt_stride = jnp.int32(1) << (m - 1 - q2)
|
|
406
|
+
is_c0 = (my_id & tgt_stride) == 0
|
|
407
|
+
new_c0 = u00 * _row + u01 * paired_row
|
|
408
|
+
new_c1 = u10 * paired_row + u11 * _row
|
|
409
|
+
after = jnp.where(is_c0, new_c0, new_c1)
|
|
410
|
+
return jnp.where(ctrl_set, after, _row)
|
|
411
|
+
|
|
412
|
+
result_2q = jnp.where(
|
|
413
|
+
q1_chunk & q2_chunk, case_2q_both_chunk(local_row),
|
|
414
|
+
jnp.where(q1_chunk & (~q2_chunk), case_2q_ctrl_chunk_tgt_local(local_row),
|
|
415
|
+
jnp.where((~q1_chunk) & q2_chunk, case_2q_ctrl_local_tgt_chunk(local_row),
|
|
416
|
+
case_2q_local_local(local_row))))
|
|
417
|
+
result_1q = jnp.where(q1_chunk, case_1q_chunk(local_row), case_1q_local(local_row))
|
|
418
|
+
|
|
419
|
+
new_row = jnp.where(is_2q, result_2q, result_1q)
|
|
420
|
+
return new_row.astype(dtype), None
|
|
421
|
+
|
|
422
|
+
return step
|
|
423
|
+
|
|
424
|
+
|
|
425
|
+
@functools.lru_cache(maxsize=None)
|
|
426
|
+
def _build_distributed_chunk_runner(num_chunks: int, m: int, k: int):
|
|
427
|
+
"""shard_map-wrapped runner: one chunk row per physical JAX device.
|
|
428
|
+
Requires jax.device_count() >= num_chunks (v1 scope: exactly one
|
|
429
|
+
chunk per device, the literal reading of issue #1 -- "distribuire
|
|
430
|
+
i chunk su più device"; a hybrid scheme with several chunks per
|
|
431
|
+
device is a possible future refinement, not attempted here).
|
|
432
|
+
|
|
433
|
+
compiled_ops (the small [g_id, q1, q2, param] sequence, identical
|
|
434
|
+
on every device) is replicated, not sharded -- P(None, None).
|
|
435
|
+
local_row is sharded along axis 0 of the (num_chunks, chunk_dim)
|
|
436
|
+
logical array, one (chunk_dim,) row per device -- P(axis_name,
|
|
437
|
+
None) on input/output, so each device's shard is that one row.
|
|
438
|
+
|
|
439
|
+
Memoized on (num_chunks, m, k) -- the three plain ints that fully
|
|
440
|
+
determine the compiled kernel. Without this, every call built a
|
|
441
|
+
brand-new Python closure and wrapped it in a fresh jax.jit, so two
|
|
442
|
+
Chunk instances with identical geometry never hit JAX's own
|
|
443
|
+
compilation cache (that cache is keyed by wrapped-function identity,
|
|
444
|
+
not by structural equality of what the closure captured) -- each one
|
|
445
|
+
silently repaid the full XLA compile cost instead of reusing the
|
|
446
|
+
other's. Caching the builder itself, not just relying on jax.jit's
|
|
447
|
+
internal cache, is what actually fixes that: the second Chunk with
|
|
448
|
+
the same (num_chunks, m, k) gets back the exact same already-jitted
|
|
449
|
+
function object."""
|
|
450
|
+
import numpy as np
|
|
451
|
+
from jax.sharding import Mesh, PartitionSpec as P
|
|
452
|
+
|
|
453
|
+
axis_name = 'chunks'
|
|
454
|
+
step = _build_distributed_chunk_step(num_chunks, m, k, axis_name)
|
|
455
|
+
|
|
456
|
+
devices = np.array(jax.devices()[:num_chunks])
|
|
457
|
+
mesh = Mesh(devices, axis_names=(axis_name,))
|
|
458
|
+
|
|
459
|
+
def run_local(local_shard, compiled_ops):
|
|
460
|
+
# shard_map keeps the sharded axis in the local shape (size
|
|
461
|
+
# num_chunks/mesh_size along axis 0 -- 1 in this v1 one-
|
|
462
|
+
# chunk-per-device scope), it doesn't squeeze it away: a
|
|
463
|
+
# (num_chunks, chunk_dim) input shards to (1, chunk_dim) per
|
|
464
|
+
# device here, not (chunk_dim,). `step` itself works on a
|
|
465
|
+
# clean (chunk_dim,) row -- squeeze going in, restore going out.
|
|
466
|
+
local_row = local_shard[0]
|
|
467
|
+
final_row, _ = jax.lax.scan(step, local_row, compiled_ops)
|
|
468
|
+
return final_row[None, :]
|
|
469
|
+
|
|
470
|
+
sharded_run = jax.shard_map(
|
|
471
|
+
run_local,
|
|
472
|
+
mesh=mesh,
|
|
473
|
+
in_specs=(P(axis_name, None), P(None, None)),
|
|
474
|
+
out_specs=P(axis_name, None),
|
|
475
|
+
check_vma=False,
|
|
476
|
+
)
|
|
477
|
+
return jax.jit(sharded_run), mesh
|
|
478
|
+
|
|
479
|
+
|
|
480
|
+
@functools.lru_cache(maxsize=None)
|
|
481
|
+
def _build_multi_chunk_runner(num_chunks: int, m: int, k: int):
|
|
482
|
+
"""jax.jit-compiled (chunks, compiled_ops) -> final_chunks, closed
|
|
483
|
+
over the static per-Chunk-instance geometry (num_chunks, m, k don't
|
|
484
|
+
change across calls on the same instance).
|
|
485
|
+
|
|
486
|
+
Memoized on (num_chunks, m, k) for the same reason
|
|
487
|
+
_build_distributed_chunk_runner is: two Chunk instances built with
|
|
488
|
+
the same geometry used to each pay a full, independent XLA compile
|
|
489
|
+
(a fresh Python closure every call means a fresh jax.jit wrapper,
|
|
490
|
+
which JAX's own compilation cache -- keyed by wrapped-function
|
|
491
|
+
identity -- can never recognize as "the same function" across
|
|
492
|
+
instances). Caching the builder means the second Chunk with matching
|
|
493
|
+
(num_chunks, m, k) gets the first one's already-compiled function
|
|
494
|
+
back directly, no recompilation at all."""
|
|
495
|
+
step = _build_multi_chunk_step(num_chunks, m, k)
|
|
496
|
+
|
|
497
|
+
@jax.jit
|
|
498
|
+
def run(chunks, compiled_ops):
|
|
499
|
+
final, _ = jax.lax.scan(step, chunks, compiled_ops)
|
|
500
|
+
return final
|
|
501
|
+
|
|
502
|
+
return run
|
|
503
|
+
|
|
504
|
+
|
|
505
|
+
def _compile_multi_chunk_ops(circuit: List) -> "jnp.ndarray":
|
|
506
|
+
"""Structural + GATE_IDS compilation shared by the multi-chunk JIT
|
|
507
|
+
path — same [g_id, q1, q2, param] row format as beast-mode's own
|
|
508
|
+
compiled ops, built via GATE_IDS instead of the old _resolve_gate's
|
|
509
|
+
GATES/PARAMETRIC_GATES lookup. This finally aligns multi-chunk's
|
|
510
|
+
gate coverage with beast-mode's (both silently skip a gate name not
|
|
511
|
+
in GATE_IDS — same known, tracked behavior, see issue #4 — instead
|
|
512
|
+
of the old _resolve_gate's NotImplementedError for e.g. ecr/iswap)."""
|
|
513
|
+
target = QuantumTranspiler.transpile(circuit)
|
|
514
|
+
rows = []
|
|
515
|
+
for cmd in target:
|
|
516
|
+
name = cmd[0].lower() if isinstance(cmd[0], str) else str(cmd[0]).lower()
|
|
517
|
+
if name not in GATE_IDS:
|
|
518
|
+
continue
|
|
519
|
+
g_id = float(GATE_IDS[name])
|
|
520
|
+
args = cmd[1:]
|
|
521
|
+
if name in ('cx', 'cz', 'cp', 'cphase', 'cy', 'crz'):
|
|
522
|
+
q1, q2 = float(args[0]), float(args[1])
|
|
523
|
+
param = float(args[2]) if len(args) > 2 else 0.0
|
|
524
|
+
rows.append([g_id, q1, q2, param])
|
|
525
|
+
elif args:
|
|
526
|
+
q1 = float(args[0])
|
|
527
|
+
param = float(args[1]) if len(args) > 1 else 0.0
|
|
528
|
+
rows.append([g_id, q1, 0.0, param])
|
|
529
|
+
if not rows:
|
|
530
|
+
return jnp.empty((0, 4), dtype=jnp.float64)
|
|
531
|
+
return jnp.array(rows, dtype=jnp.float64)
|