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,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)