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,
) -> np.ndarray:
"""Esegue le 4 fasi (scudo entrata -> modello -> scudo uscita) e ritorna l'output protetto."""
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)
for b in range(B):
flat_cl, flat_co = x_reference_np[b].flatten(), x_corrupted_np[b].flatten()
total_elements = flat_cl.size
out_flat = np.zeros_like(flat_cl)
margine_flat = np.zeros_like(flat_cl)
c_size = self.chunk_threshold if total_elements > self.chunk_threshold else total_elements
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)
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)
logger.info("Transito concluso. Sistema sigillato in %.3f secondi totali.", time.time() - t_start)
return x_final