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,399 @@
1
+ from collections import OrderedDict
2
+
3
+ import numpy as np
4
+ from scipy.ndimage import median_filter
5
+
6
+ # Issue #258 point 7: an MCP agent calling dense_evolution_vector_healing
7
+ # typically uses a different `n` (sequence length) almost every call (50,
8
+ # then 200, then 1200 steps), not the same one repeatedly -- measured
9
+ # directly: a genuinely new n costs ~0.5s (a fresh XLA compile of the
10
+ # vmapped 'phi' path below) vs ~0.013s once that exact n has recurred, a
11
+ # ~40x penalty that hits nearly every call in that usage pattern. This
12
+ # bounded LRU tracks which batch sizes (n-2, the vmap axis below) have
13
+ # already paid that compile once THIS process -- see _phi_trigger_batch's
14
+ # own branch on it.
15
+ _PHI_VMAP_WARM_SIZES: "OrderedDict[int, None]" = OrderedDict()
16
+ _PHI_VMAP_WARM_MAXSIZE = 8
17
+
18
+
19
+ def _mark_and_check_warm(size: int) -> bool:
20
+ """Returns True if `size` was already warmed (JAX's own jit cache will
21
+ reuse the compiled executable for it, so the vmapped path is free);
22
+ False the first time `size` is seen this process, in which case the
23
+ caller should use the per-index loop instead of paying a fresh XLA
24
+ compile for a batch size the MCP calling pattern likely never revisits."""
25
+ was_warm = size in _PHI_VMAP_WARM_SIZES
26
+ _PHI_VMAP_WARM_SIZES[size] = None
27
+ _PHI_VMAP_WARM_SIZES.move_to_end(size)
28
+ if len(_PHI_VMAP_WARM_SIZES) > _PHI_VMAP_WARM_MAXSIZE:
29
+ _PHI_VMAP_WARM_SIZES.popitem(last=False)
30
+ return was_warm
31
+
32
+ def median_healing(vettori: np.ndarray, radius_baseline: int = None) -> (np.ndarray, int):
33
+ """
34
+ Applica un filtro mediano avanzato ai vettori.
35
+
36
+ Questo metodo calcola un raggio per il filtro mediano dinamicamente, se non specificato,
37
+ e utilizza `scipy.ndimage.median_filter` per un'applicazione efficiente. Gestisce i
38
+ bordi della sequenza tramite padding 'nearest' e preprocessa i vettori per gestire
39
+ valori `np.nan` e `np.inf` prima dell'applicazione del filtro.
40
+
41
+ Args:
42
+ vettori (np.ndarray): Array di vettori di hidden states (n_tokens, hidden_dim).
43
+ radius_baseline (int, optional): Raggio fisso per il calcolo della mediana.
44
+ Se `None`, il raggio viene calcolato dinamicamente
45
+ come `min(20, max(3, n_tokens // 3))`.
46
+ Defaults to None.
47
+
48
+ Returns:
49
+ tuple: Contiene:
50
+ - np.ndarray: Vettori con filtro mediano applicato, della stessa shape dell'input.
51
+ - int: Il raggio effettivamente utilizzato per il filtro mediano.
52
+ """
53
+ vettori = np.asarray(vettori)
54
+ n, hidden_dim = vettori.shape
55
+
56
+ if n == 0:
57
+ return np.empty((0, hidden_dim)), 0
58
+
59
+ processed_vettori = np.copy(vettori)
60
+ processed_vettori[np.isinf(processed_vettori)] = np.nan
61
+
62
+ # A column that's entirely NaN makes np.nanmean raise "RuntimeWarning:
63
+ # Mean of empty slice" and return NaN for it (silently caught by the
64
+ # next line, which zeroes it anyway) -- pre-replacing whole all-NaN
65
+ # columns with 0.0 means nanmean never sees an empty slice, so the
66
+ # warning never fires, with byte-identical output to before.
67
+ all_nan_cols = np.all(np.isnan(processed_vettori), axis=0)
68
+ safe_for_mean = np.where(all_nan_cols, 0.0, processed_vettori)
69
+ col_means = np.nanmean(safe_for_mean, axis=0)
70
+ processed_vettori = np.where(np.isnan(processed_vettori), col_means, processed_vettori)
71
+
72
+ if n < 3:
73
+ calculated_radius = 0
74
+ window_size = 1
75
+ elif radius_baseline is None:
76
+ calculated_radius = min(20, max(3, n // 3))
77
+ window_size = 2 * calculated_radius + 1
78
+ else:
79
+ calculated_radius = radius_baseline
80
+ window_size = 2 * calculated_radius + 1
81
+
82
+ window_size = max(1, min(window_size, n))
83
+
84
+ out = median_filter(processed_vettori, size=(window_size, 1), mode='nearest')
85
+
86
+ return out, calculated_radius
87
+
88
+ def enhanced_dense_healing_hybrid(
89
+ vettori: np.ndarray,
90
+ radius_baseline: int = None,
91
+ trigger_mode: str = 'phi',
92
+ ) -> (np.ndarray, dict):
93
+ """
94
+ Applica una strategia di healing ibrida combinando la logica di dense_evolution
95
+ con un fallback alla mediana, decidendo dinamicamente quale approccio utilizzare.
96
+
97
+ Questa funzione preprocessa i vettori per gestire `np.nan` e `np.inf`.
98
+ Include telemetria dettagliata per monitorare il comportamento del processo di healing.
99
+
100
+ Args:
101
+ vettori (np.ndarray): Array di vettori di hidden states (n_tokens, hidden_dim).
102
+ radius_baseline (int, optional): Raggio fisso per il calcolo delle baseline (media/mediana).
103
+ Se `None`, il raggio viene calcolato dinamicamente
104
+ come `min(20, max(3, n_tokens // 3))`.
105
+ Defaults to None.
106
+ trigger_mode (str, optional): Quale meccanismo decide se un dato passo è
107
+ movimento genuino (mantenuto com'è) o
108
+ rumore/corruzione (sostituito con la mediana
109
+ locale). Uno tra:
110
+ - 'phi' (default): il Phi-Trigger originale
111
+ (dense_evolution.mitigation.healing.evaluate_phi_trigger),
112
+ soglia fissa |v_dinamic| > 0.01. Mantenuto
113
+ come default per piena compatibilità
114
+ all'indietro -- è anche il meccanismo esatto
115
+ preso di mira dal red-teaming a gradiente di
116
+ ia_utils.adversarial_vector_attack, dato che
117
+ calculate_phi_ab/calculate_vettore_dinamico
118
+ sono funzioni JAX differenziabili.
119
+ - 'adaptive': trigger a deviazione locale
120
+ adattiva (MAD), consapevole di NaN/Inf
121
+ (Dense-Evolution-Discovery Esperimento 27).
122
+ Validato: riduce il tasso di falsi positivi
123
+ (sostituzioni su dati rumorosi ma non
124
+ corrotti) dall'~90% al ~12%, mantenendo un
125
+ tasso di rilevamento delle corruzioni reali
126
+ pari o superiore al Phi-Trigger su ogni tipo
127
+ testato (picchi singoli, sequenze di NaN,
128
+ outlier sparsi, corruzioni combinate). Non
129
+ differenziabile (usa np.median/np.std), quindi
130
+ il red-teaming a gradiente non si applica
131
+ allo stesso modo.
132
+ Defaults to 'phi'.
133
+
134
+ Returns:
135
+ tuple: Contiene:
136
+ - np.ndarray: Vettori curati, della stessa shape dell'input.
137
+ - dict: Metadati di telemetria contenenti:
138
+ - 'fallback_triggered' (bool): `True` solo se l'input originale conteneva
139
+ NaN/Inf E il fallback mediano è stato applicato
140
+ almeno una volta per correggerlo. Non riflette
141
+ correzioni del trigger su dati validi ma
142
+ "staticamente" rumorosi (nessuna corruzione reale).
143
+ - 'adaptive_radius_used' (int): Il raggio effettivamente calcolato e applicato.
144
+ - 'reconstruction_error' (float): La norma media di variazione (errore di ricostruzione)
145
+ introdotta rispetto ai vettori originali (potenzialmente corrotti).
146
+ - 'trigger_mode' (str): Il meccanismo di trigger effettivamente usato.
147
+ """
148
+ if trigger_mode not in ('phi', 'adaptive'):
149
+ raise ValueError(f"trigger_mode must be 'phi' or 'adaptive', got {trigger_mode!r}")
150
+
151
+ if trigger_mode == 'phi':
152
+ try:
153
+ import jax
154
+ import jax.numpy as jnp
155
+ from dense_evolution.mitigation.healing import (
156
+ calculate_phi_ab,
157
+ calculate_vettore_dinamico,
158
+ evaluate_phi_trigger,
159
+ GLOBAL_CONSTANTS,
160
+ )
161
+ except ImportError as _import_error:
162
+ # jax is a core dependency of dense-evolution (see pyproject.toml),
163
+ # so this shouldn't fail on a normal install -- but ia_utils could
164
+ # be used standalone outside the full package (e.g. a stripped or
165
+ # vendored copy missing dense_evolution.healing, or an environment
166
+ # missing jax), where the bare ModuleNotFoundError gives no hint
167
+ # this module needs them.
168
+ raise ImportError(
169
+ "enhanced_dense_healing_hybrid requires jax and dense_evolution.healing "
170
+ f"(failed to import: {_import_error}). Install the full dense-evolution "
171
+ "package (which depends on jax) to use this function."
172
+ ) from _import_error
173
+
174
+ n, hidden_dim = vettori.shape
175
+
176
+ if n == 0:
177
+ return np.empty((0, hidden_dim)), {'fallback_triggered': False, 'adaptive_radius_used': 0,
178
+ 'reconstruction_error': 0.0, 'trigger_mode': trigger_mode}
179
+
180
+ # Computed on the RAW input, before any sanitization -- fallback_triggered
181
+ # in the returned metadata is gated on this (see below), so it reflects
182
+ # "there was genuine NaN/Inf corruption AND the median fallback fired",
183
+ # not just "the trigger's internal heuristic called some row static".
184
+ # The 'phi' heuristic alone also fires on structurally noisy-but-valid
185
+ # data (e.g. pure IID random input with no coherent trend for it to
186
+ # recognize as genuine motion) -- verified directly: clean random
187
+ # Gaussian input with zero NaN/Inf still tripped the un-gated flag.
188
+ had_nan_or_inf = bool(np.isnan(vettori).any() or np.isinf(vettori).any())
189
+ # Per-row raw corruption flag ('adaptive' mode only): forces healing at
190
+ # any row that was originally NaN/Inf, regardless of the deviation
191
+ # statistic -- after NaN/Inf sanitization below, a corrupted row is
192
+ # replaced by the (column-wise) global mean, which can look
193
+ # statistically unremarkable relative to the LOCAL window and evade a
194
+ # purely deviation-based trigger. Verified in Discovery Experiment 27:
195
+ # the adaptive trigger without this row-level override missed 100% of
196
+ # NaN-run corruption for exactly this reason.
197
+ raw_nan_or_inf_row = np.isnan(vettori).any(axis=1) | np.isinf(vettori).any(axis=1)
198
+
199
+ processed_vettori = np.copy(vettori)
200
+ processed_vettori[np.isinf(processed_vettori)] = np.nan
201
+
202
+ # See median_healing's identical block above for why this avoids
203
+ # np.nanmean's "Mean of empty slice" warning on an all-NaN column.
204
+ all_nan_cols = np.all(np.isnan(processed_vettori), axis=0)
205
+ safe_for_mean = np.where(all_nan_cols, 0.0, processed_vettori)
206
+ col_means = np.nanmean(safe_for_mean, axis=0)
207
+ processed_vettori = np.where(np.isnan(processed_vettori), col_means, processed_vettori)
208
+
209
+ out = np.copy(processed_vettori)
210
+
211
+ if radius_baseline is None:
212
+ if n < 3:
213
+ adaptive_radius_used = 0
214
+ else:
215
+ adaptive_radius_used = min(20, max(3, n // 3))
216
+ else:
217
+ adaptive_radius_used = radius_baseline
218
+
219
+ fallback_triggered_at_all = False
220
+ reconstruction_errors_per_step = []
221
+
222
+ if n > 0:
223
+ reconstruction_errors_per_step.append(np.linalg.norm(out[0] - processed_vettori[0]))
224
+ if n > 1:
225
+ reconstruction_errors_per_step.append(np.linalg.norm(out[1] - processed_vettori[1]))
226
+
227
+ if trigger_mode == 'phi':
228
+ # BUG FIX (perf, prog.txt Sezione 4.2): the old loop converted
229
+ # baseline_mean/processed_vettori[i]/ipg_vector to jnp.array and the
230
+ # trigger decision back to a Python float on EVERY iteration -- a
231
+ # NumPy<->JAX round trip per step, the exact cost jax.lax.scan/vmap
232
+ # exist to avoid. calculate_phi_ab/calculate_vettore_dinamico/
233
+ # evaluate_phi_trigger take fixed-shape (hidden_dim,) or scalar
234
+ # inputs (no variable-size window inside them -- only the median
235
+ # fallback below has one, and that stays plain NumPy, unchanged),
236
+ # so the whole per-index trigger computation batches cleanly with
237
+ # jax.vmap: ONE host<->device round trip for the whole sequence
238
+ # instead of one per index. Values, not just speed, are unchanged --
239
+ # every array below is built with the identical NumPy arithmetic
240
+ # the old loop did per-step, just evaluated for every index at once.
241
+ # Measured (n, hidden_dim=32, same-shape warm call both versions):
242
+ # n=50 1.6x, n=300 5.0x, n=1000 9.8x, n=3000 11.1x faster than the
243
+ # old per-step loop. Real trade-off, not hidden: the old loop's
244
+ # jax.jit'd calls operate on fixed (hidden_dim,) shapes, so JAX
245
+ # compiles them ONCE ever for a given hidden_dim, reused across
246
+ # every n. This vmapped version's batch axis is n-2, so a NEW n
247
+ # means a fresh XLA compile the first time that exact length is
248
+ # seen -- worth it whenever the same length recurs (the normal
249
+ # case: one call per input sequence, not a different n each time),
250
+ # not free on a single one-off call at a brand new length.
251
+ idx = np.arange(2, n)
252
+ if idx.size:
253
+ radius = adaptive_radius_used
254
+ lo_arr = np.maximum(0, idx - radius)
255
+
256
+ # baseline_mean[k] = mean(processed_vettori[lo_arr[k]:idx[k]]) via
257
+ # a prefix sum -- same value the old sliding-window-sum loop
258
+ # computed per step, done here for every index in one shot.
259
+ prefix_sum = np.concatenate(
260
+ [np.zeros((1, hidden_dim)), np.cumsum(processed_vettori, axis=0)], axis=0
261
+ )
262
+ window_counts = (idx - lo_arr).astype(np.float64)
263
+ baseline_means = (prefix_sum[idx] - prefix_sum[lo_arr]) / window_counts[:, None]
264
+
265
+ ipg_raw = processed_vettori[idx - 1] - processed_vettori[idx - 2]
266
+ norm_ipg_raw = np.linalg.norm(ipg_raw, axis=1)
267
+ safe_norm = np.where(norm_ipg_raw > 1e-9, norm_ipg_raw, 1.0)
268
+ ipg_vectors = np.where((norm_ipg_raw > 1e-9)[:, None], ipg_raw / safe_norm[:, None], ipg_raw)
269
+
270
+ if _mark_and_check_warm(idx.size):
271
+ # This batch size has been compiled once already this
272
+ # process -- JAX's own jit cache (keyed on calculate_phi_ab
273
+ # et al.'s abstract input shapes, batch dim included) reuses
274
+ # that executable, so vmap really is ~free here.
275
+ state_A = jnp.asarray(baseline_means)
276
+ state_B = jnp.asarray(processed_vettori[idx])
277
+ ipg_vector_batch = jnp.asarray(ipg_vectors)
278
+
279
+ phi_ab = jax.vmap(calculate_phi_ab)(state_A, state_B, ipg_vector_batch)
280
+ E_A = jnp.linalg.norm(state_A, axis=1)
281
+ E_B = jnp.linalg.norm(state_B, axis=1)
282
+ v_dinamic = jax.vmap(calculate_vettore_dinamico)(E_A, E_B, phi_ab)
283
+ trigger, _, _ = jax.vmap(evaluate_phi_trigger)(v_dinamic)
284
+ is_dynamic_arr = np.asarray(trigger) > GLOBAL_CONSTANTS['NON_STATIC_THRESHOLD_A']
285
+ else:
286
+ # First time this exact batch size is seen this process --
287
+ # calculate_phi_ab/calculate_vettore_dinamico/evaluate_phi_trigger
288
+ # are each jitted on a fixed (hidden_dim,)/scalar shape,
289
+ # compiled once ever regardless of n, so calling them one
290
+ # index at a time here pays no compile cost at all (unlike
291
+ # vmap's batch-shaped call above, which would). Same
292
+ # arithmetic as the batched path, just per-index -- verified
293
+ # to match it exactly in tests/unit/test_ia_utils_vector_healing.py.
294
+ is_dynamic_list = []
295
+ for k in range(idx.size):
296
+ state_A_k = jnp.asarray(baseline_means[k])
297
+ state_B_k = jnp.asarray(processed_vettori[idx[k]])
298
+ ipg_k = jnp.asarray(ipg_vectors[k])
299
+ phi_ab_k = calculate_phi_ab(state_A_k, state_B_k, ipg_k)
300
+ v_dinamic_k = calculate_vettore_dinamico(
301
+ jnp.linalg.norm(state_A_k), jnp.linalg.norm(state_B_k), phi_ab_k
302
+ )
303
+ trigger_k, _, _ = evaluate_phi_trigger(v_dinamic_k)
304
+ is_dynamic_list.append(bool(np.asarray(trigger_k) > GLOBAL_CONSTANTS['NON_STATIC_THRESHOLD_A']))
305
+ is_dynamic_arr = np.array(is_dynamic_list, dtype=bool)
306
+
307
+ # trigger == 1.0: ciclo aperto/dinamico -> cambio genuino, si tiene il valore
308
+ # trigger == 0.0: ciclo chiuso/statico -> rumore, si sostituisce con la mediana locale
309
+
310
+ dynamic_idx = idx[is_dynamic_arr]
311
+ out[dynamic_idx] = processed_vettori[dynamic_idx]
312
+
313
+ # The median fallback's window has a genuinely variable size
314
+ # (grows from 1 up to radius) -- not batchable the same way,
315
+ # but it only runs for the (typically minority) indices the
316
+ # trigger actually flags as noise, same np.median as before.
317
+ for i, lo in zip(idx[~is_dynamic_arr], lo_arr[~is_dynamic_arr]):
318
+ out[i] = np.median(processed_vettori[lo:i], axis=0)
319
+ fallback_triggered_at_all = True
320
+
321
+ reconstruction_errors_per_step.extend(
322
+ np.linalg.norm(out[idx] - processed_vettori[idx], axis=1).tolist()
323
+ )
324
+
325
+ else: # trigger_mode == 'adaptive'
326
+ # BUG FIX (perf): baseline_mean used to be np.mean(processed_vettori[lo:i])
327
+ # recomputed from scratch every iteration -- O(window_size) per step.
328
+ # window_size is capped at 20 for the default adaptive radius, but an
329
+ # explicit radius_baseline (a caller-supplied parameter, unbounded) can
330
+ # make it grow with i, making the whole loop O(n^2) in that case. A
331
+ # sliding-window sum (add the newly-entering element, subtract the
332
+ # element that just fell out of the window) makes each step O(1)
333
+ # amortized regardless of radius_baseline, at the one-time cost of a
334
+ # single O(window_size) sum for the first window.
335
+ #
336
+ # Deliberately NOT converted to JAX like the 'phi' branch above:
337
+ # this trigger uses np.median/np.std on purpose, precisely so it is
338
+ # NOT part of the JAX-differentiable graph gradient-based red-teaming
339
+ # (ia_utils.adversarial_vector_attack) can attack -- see this
340
+ # function's own docstring. Vectorizing it via JAX would silently
341
+ # remove that property, not just speed things up.
342
+ window_sum = np.sum(processed_vettori[max(0, 2 - adaptive_radius_used):2], axis=0)
343
+ window_lo = max(0, 2 - adaptive_radius_used)
344
+
345
+ # Running history of each step's own local deviation (norm from its
346
+ # window mean, normalized by sqrt(hidden_dim)), reused as the
347
+ # recent-deviation sample for later steps' adaptive threshold
348
+ # instead of recomputing it from scratch each time.
349
+ deviation_history = []
350
+
351
+ for i in range(2, n):
352
+ lo = max(0, i - adaptive_radius_used)
353
+ if i > 2:
354
+ window_sum = window_sum + processed_vettori[i - 1]
355
+ for dropped_idx in range(window_lo, lo):
356
+ window_sum = window_sum - processed_vettori[dropped_idx]
357
+ window_lo = lo
358
+ baseline_mean = window_sum / (i - lo)
359
+
360
+ current_deviation = np.linalg.norm(processed_vettori[i] - baseline_mean) / np.sqrt(hidden_dim)
361
+
362
+ recent = deviation_history[-adaptive_radius_used:] if adaptive_radius_used > 0 else []
363
+ if len(recent) > 1:
364
+ recent_arr = np.array(recent)
365
+ local_median = np.median(recent_arr)
366
+ # Median Absolute Deviation, scaled by 1.4826 to be a
367
+ # consistent estimator of std under normality -- robust to
368
+ # the very outliers it is meant to help detect (a raw std
369
+ # over a window containing an outlier is itself inflated by
370
+ # that outlier, which is exactly what let scattered-outlier
371
+ # corruption partially evade a std-based version of this
372
+ # trigger in Discovery Experiment 27).
373
+ local_spread = 1.4826 * np.median(np.abs(recent_arr - local_median))
374
+ else:
375
+ local_median, local_spread = 0.1, 0.05
376
+ adaptive_threshold = max(local_median + 3.5 * local_spread, 0.25)
377
+
378
+ is_dynamic = (current_deviation < adaptive_threshold) and not raw_nan_or_inf_row[i]
379
+ deviation_history.append(current_deviation)
380
+
381
+ if is_dynamic:
382
+ healed_vector = processed_vettori[i]
383
+ else:
384
+ healed_vector = np.median(processed_vettori[lo:i], axis=0)
385
+ fallback_triggered_at_all = True
386
+
387
+ out[i] = healed_vector
388
+ reconstruction_errors_per_step.append(np.linalg.norm(out[i] - processed_vettori[i]))
389
+
390
+ mean_reconstruction_error = np.mean(reconstruction_errors_per_step) if reconstruction_errors_per_step else 0.0
391
+
392
+ metadata = {
393
+ 'fallback_triggered': fallback_triggered_at_all and had_nan_or_inf,
394
+ 'adaptive_radius_used': adaptive_radius_used,
395
+ 'reconstruction_error': mean_reconstruction_error,
396
+ 'trigger_mode': trigger_mode,
397
+ }
398
+
399
+ return out, metadata
local_site/__init__.py ADDED
File without changes
File without changes