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