def protect_and_forward(
self,
ai_model_callable: Optional[Callable],
x_corrupted: np.ndarray,
x_reference: Optional[np.ndarray] = None,
use_input_shield: bool = True,
use_model_injection: bool = True,
use_output_shield: bool = True,
use_arbiter: bool = False,
) -> np.ndarray:
"""Esegue le 4 fasi (scudo entrata -> modello -> scudo uscita) e ritorna l'output protetto.
use_arbiter — se True, dopo lo scudo entrata standard instrada ogni
punto verso il correttore giusto (utility/arbiter.py): impulso
isolato -> rigetto duro, cambio di regime sostenuto -> passa
grezzo, nessuna anomalia puntuale -> resta il risultato gia'
prodotto dallo scudo standard (il suo smorzamento morbido, non un
pass-through). Default False: comportamento e risultati identici a
prima di questa opzione su tutti i test esistenti. Popola
etichette_arbitro/incertezza_arbitro; se x_reference e' noto,
aggiorna anche tipi_corruzione_visti(slice_shape)."""
is_simple_data_test = ai_model_callable is None or not use_model_injection
if is_simple_data_test:
logger.info("CONTRAZIONE LOGICA DETECTED: Riconosciuto Test di Protezione Dati Semplice (No IA Model).")
use_model_injection = False
# Un array 1D (es. una singola serie da sensore/pipeline) NON e' un
# batch di N scalari indipendenti: e' UNA istanza con N campioni
# correlati nel tempo. Senza questa promozione, B=N e ogni campione
# veniva processato da solo (slice_shape=()), azzerando il contesto
# su cui si basa il rilevamento outlier in modalita' cieca (senza
# x_reference) -- il caso d'uso principale documentato nel README.
was_1d = (x_corrupted.ndim == 1)
if was_1d:
x_corrupted = np.asarray(x_corrupted).reshape(1, -1)
if x_reference is not None:
x_reference = np.asarray(x_reference).reshape(1, -1)
orig_shape = x_corrupted.shape
B = orig_shape[0]
slice_shape = orig_shape[1:]
t_start = time.time()
if use_input_shield:
logger.info("Attivazione SCUDO ENTRATA (4 Fasi) su Ipervolume: %s", orig_shape)
x_corrupted_np = np.array(x_corrupted)
if x_reference is None:
x_reference_np = np.zeros_like(x_corrupted_np)
for b in range(B):
row_flat = x_corrupted_np[b].flatten()
recalled = self._recall_reference(row_flat, slice_shape)
if recalled is not None:
x_reference_np[b] = recalled.reshape(slice_shape)
else:
x_reference_np[b] = self._blind_reference(row_flat).reshape(slice_shape)
else:
x_reference_np = np.array(x_reference)
self._remember_reference(x_reference_np, slice_shape)
purified_batch = np.zeros(orig_shape, dtype=np.float64)
margine_batch = np.zeros(orig_shape, dtype=np.float64)
flat_cl_all = x_reference_np.reshape(B, -1)
flat_co_all = x_corrupted_np.reshape(B, -1)
total_elements = flat_cl_all.shape[1]
if total_elements <= self.chunk_threshold:
# Percorso vettorizzato: le B righe sono processate in UNA
# chiamata (jax.vmap sotto, calibrazione indipendente per
# riga) invece di B dispatch JAX separati -- vedi
# _execute_4_phase_input_shield_batch e CHANGELOG per il
# guadagno misurato. Ogni riga e' completamente indipendente
# dalle altre a questo stadio (nessuno stato condiviso tra
# righe qui: _reference_bank/_corruption_type_memory non
# sono toccati in questo blocco).
out_batch, marg_batch = self._execute_4_phase_input_shield_batch(flat_cl_all, flat_co_all)
purified_batch = out_batch.reshape(orig_shape)
margine_batch = marg_batch.reshape(orig_shape)
self._gc_se_ram_bassa()
else:
# Fallback invariato: una singola riga eccede da sola
# chunk_threshold (raro), va comunque spezzata in blocchi.
for b in range(B):
flat_cl, flat_co = flat_cl_all[b], flat_co_all[b]
out_flat = np.zeros_like(flat_cl)
margine_flat = np.zeros_like(flat_cl)
c_size = self.chunk_threshold
for start_idx in range(0, total_elements, c_size):
end_idx = min(start_idx + c_size, total_elements)
purified_chunk, margine_chunk = self._execute_4_phase_input_shield(flat_cl[start_idx:end_idx], flat_co[start_idx:end_idx])
out_flat[start_idx:end_idx] = purified_chunk
margine_flat[start_idx:end_idx] = margine_chunk
self._gc_se_ram_bassa()
purified_batch[b] = out_flat.reshape(slice_shape)
margine_batch[b] = margine_flat.reshape(slice_shape)
if use_arbiter:
etichette_batch = np.empty(orig_shape, dtype=object)
incertezza_batch = np.zeros(orig_shape, dtype=np.float64)
for b in range(B):
co_flat = x_corrupted_np[b].flatten()
pur_flat = purified_batch[b].flatten()
corretto_flat, etichette_flat, incertezza_flat = self._applica_arbitro(co_flat, pur_flat)
purified_batch[b] = corretto_flat.reshape(slice_shape)
etichette_batch[b] = etichette_flat.reshape(slice_shape)
incertezza_batch[b] = incertezza_flat.reshape(slice_shape)
if x_reference is not None:
self._aggiorna_memoria_corruzione(etichette_flat, slice_shape)
self.etichette_arbitro = etichette_batch
self.incertezza_arbitro = incertezza_batch
self.incertezza_arbitro_media = float(np.mean(incertezza_batch))
x_for_model = jnp.array(purified_batch)
self.margine_ingresso = margine_batch
self.margine_ingresso_medio = float(np.mean(margine_batch))
self.margine_ingresso_max = float(np.max(margine_batch))
logger.info("Input purificato in %.3fs. Margine d'errore: medio=%.4g, max=%.4g",
time.time() - t_start, self.margine_ingresso_medio, self.margine_ingresso_max)
else:
logger.info("SCUDO ENTRATA disattivato. I dati transitano senza pre-filtri.")
x_for_model = jnp.array(x_corrupted)
self.margine_ingresso = None
self.margine_ingresso_medio = self.margine_ingresso_max = 0.0
if use_model_injection:
t_ia = time.time()
ai_output = ai_model_callable(x_for_model)
jax.block_until_ready(ai_output)
logger.info("Risposta IA ottenuta in %.3fs. Shape Output: %s", time.time() - t_ia, ai_output.shape)
else:
logger.info("INIEZIONE MODELLO bypassata. I dati purificati procedono verso la barriera spettrale.")
ai_output = x_for_model
if use_output_shield:
t_out = time.time()
logger.info("Attivazione SCUDO USCITA (4 Fasi) su Spettro Terminale...")
if use_model_injection and x_reference is not None:
# riferimento nello SPAZIO DI OUTPUT: risposta del modello al dato
# pulito, non l'input purificato (spazio diverso se il modello e'
# trasformativo, es. classificatori/embedding/reti non-lineari)
output_reference = ai_model_callable(jnp.array(x_reference))
jax.block_until_ready(output_reference)
else:
# nessun riferimento pulito noto: auto-consistenza cieca sull'output
# stesso (stesso stabilizzatore della Fase 1, applicato qui all'uscita)
output_reference = self.stabilizer.filter_batch_scenarios(
ai_output.reshape(ai_output.shape[0], -1)).reshape(ai_output.shape)
x_final, margine_out = self._execute_4_phase_output_shield(ai_output, output_reference)
self.margine_uscita = np.array(margine_out)
self.margine_uscita_medio = float(jnp.mean(margine_out))
self.margine_uscita_max = float(jnp.max(margine_out))
logger.info("Output rinormalizzato in %.3fs. Margine d'errore: medio=%.4g, max=%.4g",
time.time() - t_out, self.margine_uscita_medio, self.margine_uscita_max)
else:
logger.info("SCUDO USCITA disattivato. Emissione del flusso lineare.")
x_final = ai_output
self.margine_uscita = None
self.margine_uscita_medio = self.margine_uscita_max = 0.0
if was_1d:
x_final = x_final.reshape(-1)
if self.margine_ingresso is not None:
self.margine_ingresso = self.margine_ingresso.reshape(-1)
if self.margine_uscita is not None:
self.margine_uscita = np.asarray(self.margine_uscita).reshape(-1)
if self.etichette_arbitro is not None:
self.etichette_arbitro = self.etichette_arbitro.reshape(-1)
if self.incertezza_arbitro is not None:
self.incertezza_arbitro = self.incertezza_arbitro.reshape(-1)
logger.info("Transito concluso. Sistema sigillato in %.3f secondi totali.", time.time() - t_start)
return x_final