Cybersecurity-Projects/PROJECTS/advanced/ai-threat-detection/backend/ml/train_autoencoder.py

116 lines
3.5 KiB
Python

"""
©AngelaMos | 2026
train_autoencoder.py
"""
from typing import Any
import numpy as np
import torch
from torch.utils.data import DataLoader, TensorDataset
from ml.autoencoder import ThreatAutoencoder
from ml.scaler import FeatureScaler
def train_autoencoder(
X_normal: np.ndarray,
epochs: int = 100,
batch_size: int = 256,
lr: float = 1e-3,
percentile: float = 99.5,
val_split: float = 0.15,
patience: int = 10,
) -> dict[str, Any]:
"""
Train the autoencoder on normal-only traffic and calibrate the
anomaly detection threshold.
"""
input_dim = X_normal.shape[1]
split_idx = int(len(X_normal) * (1 - val_split))
X_train_raw = X_normal[:split_idx]
X_val_raw = X_normal[split_idx:]
scaler = FeatureScaler()
X_train_scaled = scaler.fit_transform(X_train_raw)
X_val_scaled = scaler.transform(X_val_raw)
train_tensor = torch.from_numpy(X_train_scaled)
val_tensor = torch.from_numpy(X_val_scaled)
train_loader = DataLoader(
TensorDataset(train_tensor),
batch_size=batch_size,
shuffle=True,
drop_last=len(train_tensor) > batch_size,
)
model = ThreatAutoencoder(input_dim=input_dim)
optimizer = torch.optim.AdamW(model.parameters(),
lr=lr,
weight_decay=1e-5,
betas=(0.9, 0.999))
scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer,
mode="min",
factor=0.5,
patience=5,
min_lr=1e-6)
history: dict[str, list[float]] = {"train_loss": [], "val_loss": []}
best_val_loss = float("inf")
best_state = None
epochs_without_improvement = 0
for _epoch in range(epochs):
model.train()
epoch_loss = 0.0
n_batches = 0
for (batch, ) in train_loader:
reconstructed = model(batch)
loss = torch.nn.functional.mse_loss(reconstructed, batch)
optimizer.zero_grad()
loss.backward() # type: ignore[no-untyped-call]
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
optimizer.step()
epoch_loss += loss.item()
n_batches += 1
avg_train_loss = epoch_loss / max(n_batches, 1)
history["train_loss"].append(avg_train_loss)
model.eval()
with torch.no_grad():
val_reconstructed = model(val_tensor)
val_loss = torch.nn.functional.mse_loss(val_reconstructed,
val_tensor).item()
history["val_loss"].append(val_loss)
scheduler.step(val_loss)
if val_loss < best_val_loss:
best_val_loss = val_loss
best_state = {k: v.clone() for k, v in model.state_dict().items()}
epochs_without_improvement = 0
else:
epochs_without_improvement += 1
if epochs_without_improvement >= patience:
break
if best_state is not None:
model.load_state_dict(best_state)
model.eval()
with torch.no_grad():
val_errors = model.compute_reconstruction_error(val_tensor)
threshold = float(np.percentile(val_errors.numpy(), percentile))
return {
"model": model,
"scaler": scaler,
"threshold": threshold,
"history": history,
}