| """ |
| train.py |
| ======== |
| Training loop for the LSTM-Autoencoder. |
| |
| Trains ONLY on normal traffic sessions. |
| After training, calculates a dynamic anomaly threshold from the reconstruction error distribution. |
| |
| Usage |
| ----- |
| python src/train.py --dataset csic2010 |
| python src/train.py --dataset cicids2018 |
| python src/train.py --dataset unsw |
| """ |
|
|
| import argparse |
| import json |
| import logging |
| import random |
| from pathlib import Path |
|
|
| import numpy as np |
| import torch |
| import torch.nn as nn |
| from torch.utils.data import DataLoader, TensorDataset |
|
|
| from model import ( |
| build_model_cicids2018, |
| build_model_csic2010, |
| build_model_unsw, |
| ) |
|
|
|
|
| def set_seed(seed=42): |
| random.seed(seed) |
| np.random.seed(seed) |
| torch.manual_seed(seed) |
| torch.backends.cudnn.deterministic = True |
|
|
|
|
| logging.basicConfig( |
| level=logging.INFO, |
| format="%(asctime)s %(levelname)s %(message)s", |
| datefmt="%H:%M:%S", |
| ) |
| log = logging.getLogger(__name__) |
|
|
|
|
| def train( |
| model: nn.Module, |
| X_train: np.ndarray, |
| dataset: str, |
| run_id: str | None = None, |
| epochs: int = 30, |
| batch_size: int = 256, |
| lr: float = 1e-3, |
| patience: int = 5, |
| device: str = "cpu", |
| ) -> list[float]: |
| """ |
| Train the autoencoder on normal-traffic sessions only. |
| |
| Uses MSE loss between input and reconstruction. |
| Stops early if validation loss stops improving. |
| |
| run_id : filename tag for the checkpoint, e.g. "csic2010_w5". |
| Defaults to `dataset` for backward compatibility if not given. |
| |
| Returns list of training losses per epoch. |
| """ |
| if run_id is None: |
| run_id = dataset |
|
|
| model = model.to(device) |
| optimizer = torch.optim.Adam(model.parameters(), lr=lr) |
| criterion = nn.MSELoss() |
|
|
| |
| split = int(len(X_train) * 0.9) |
| X_tr = X_train[:split] |
| X_val = X_train[split:] |
|
|
| |
| if dataset == "csic2010": |
| tr_tensor = torch.tensor(X_tr, dtype=torch.long) |
| val_tensor = torch.tensor(X_val, dtype=torch.long) |
| else: |
| X_tr = np.nan_to_num(X_tr, nan=0.0, posinf=0.0, neginf=0.0) |
| X_val = np.nan_to_num(X_val, nan=0.0, posinf=0.0, neginf=0.0) |
| tr_tensor = torch.tensor(X_tr, dtype=torch.float32) |
| val_tensor = torch.tensor(X_val, dtype=torch.float32) |
|
|
| tr_loader = DataLoader( |
| TensorDataset(tr_tensor), batch_size=batch_size, shuffle=True |
| ) |
| val_loader = DataLoader(TensorDataset(val_tensor), batch_size=batch_size) |
|
|
| train_losses = [] |
| val_losses = [] |
| best_val = float("inf") |
| patience_ctr = 0 |
|
|
| for epoch in range(1, epochs + 1): |
| |
| model.train() |
| epoch_loss = 0.0 |
|
|
| for (batch,) in tr_loader: |
| batch = batch.to(device) |
| optimizer.zero_grad() |
|
|
| recon = model(batch) |
|
|
| |
| if dataset == "csic2010": |
| target = model.embedding(batch).detach() |
| else: |
| target = batch |
|
|
| loss = criterion(recon, target) |
| loss.backward() |
|
|
| |
| nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) |
|
|
| optimizer.step() |
| epoch_loss += loss.item() * len(batch) |
|
|
| epoch_loss /= len(X_tr) |
|
|
| |
| model.eval() |
| val_loss = 0.0 |
| with torch.no_grad(): |
| for (batch,) in val_loader: |
| batch = batch.to(device) |
| recon = model(batch) |
| if dataset == "csic2010": |
| target = model.embedding(batch).detach() |
| else: |
| target = batch |
| val_loss += criterion(recon, target).item() * len(batch) |
| val_loss /= len(X_val) |
|
|
| train_losses.append(epoch_loss) |
| val_losses.append(val_loss) |
|
|
| log.info( |
| "Epoch %02d/%02d train=%.6f val=%.6f", epoch, epochs, epoch_loss, val_loss |
| ) |
|
|
| |
| if val_loss < best_val - 1e-6: |
| best_val = val_loss |
| patience_ctr = 0 |
| |
| torch.save(model.state_dict(), f"models/best_{run_id}.pt") |
| else: |
| patience_ctr += 1 |
| if patience_ctr >= patience: |
| log.info("Early stopping at epoch %d", epoch) |
| break |
|
|
| return train_losses, val_losses |
|
|
|
|
| def calculate_threshold( |
| model: nn.Module, |
| X_train: np.ndarray, |
| dataset: str, |
| percentile: float = 95.0, |
| device: str = "cpu", |
| ) -> tuple[float, float, float]: |
| """ |
| Calculate the anomaly detection threshold from normal traffic reconstruction errors. |
| |
| Strategy: fit threshold at the Nth percentile of normal errors. Anything above this is flagged as anomalous. |
| |
| Returns (threshold, mean_error, std_error) |
| """ |
| model.eval() |
| model = model.to(device) |
|
|
| if dataset == "csic2010": |
| tensor = torch.tensor(X_train, dtype=torch.long) |
| else: |
| X_train = np.nan_to_num(X_train, nan=0.0, posinf=0.0, neginf=0.0) |
| tensor = torch.tensor(X_train, dtype=torch.float32) |
|
|
| loader = DataLoader(TensorDataset(tensor), batch_size=512) |
|
|
| all_errors = [] |
| with torch.no_grad(): |
| for (batch,) in loader: |
| batch = batch.to(device) |
| errors = model.reconstruction_error(batch) |
| all_errors.extend(errors.cpu().numpy()) |
|
|
| all_errors = np.array(all_errors) |
| threshold = float(np.percentile(all_errors, percentile)) |
| mean_err = float(all_errors.mean()) |
| std_err = float(all_errors.std()) |
|
|
| log.info("Threshold (%.0fth percentile): %.6f", percentile, threshold) |
| log.info("Normal error — mean: %.6f std: %.6f", mean_err, std_err) |
|
|
| return threshold, mean_err, std_err |
|
|
|
|
| def main(): |
| set_seed(42) |
|
|
| parser = argparse.ArgumentParser() |
| parser.add_argument( |
| "--dataset", required=True, choices=["csic2010", "cicids2018", "unsw"] |
| ) |
| parser.add_argument("--epochs", type=int, default=30) |
| parser.add_argument("--batch_size", type=int, default=256) |
| parser.add_argument("--lr", type=float, default=1e-3) |
| parser.add_argument("--hidden", type=int, default=64) |
| parser.add_argument("--layers", type=int, default=2) |
| parser.add_argument("--percentile", type=float, default=95.0) |
| parser.add_argument( |
| "--window", |
| type=int, |
| default=5, |
| help="Sliding window size (must match preprocessing.py --window)", |
| ) |
| args = parser.parse_args() |
|
|
| device = "cuda" if torch.cuda.is_available() else "cpu" |
| data_dir = Path("data/processed") |
| model_dir = Path("models") |
| model_dir.mkdir(exist_ok=True) |
|
|
| run_id = f"{args.dataset}_w{args.window}" |
|
|
| log.info("Device: %s", device) |
| log.info("Dataset: %s Window: %d", args.dataset, args.window) |
|
|
| |
| X_train = np.load(data_dir / f"X_train_{run_id}.npy") |
| log.info("X_train shape: %s", X_train.shape) |
|
|
| |
| |
| |
| if args.dataset == "csic2010": |
| vocab_path = data_dir / f"vocab_{run_id}.json" |
| vocab_data = json.load(open(vocab_path)) |
| saved_window = vocab_data.get("window_size") |
| vocab = vocab_data.get("vocab", vocab_data) |
| if saved_window is not None and saved_window != args.window: |
| raise ValueError( |
| f"Window mismatch: {vocab_path.name} was built with " |
| f"window={saved_window}, but --window={args.window} was passed. " |
| f"Re-run preprocessing.py --window {args.window}, or pass " |
| f"--window {saved_window} here to match it." |
| ) |
| model = build_model_csic2010( |
| vocab_size=len(vocab), |
| hidden_size=args.hidden, |
| num_layers=args.layers, |
| seq_len=args.window, |
| ) |
| elif args.dataset == "cicids2018": |
| scaler_path = data_dir / f"scaler_{run_id}.json" |
| scaler_data = json.load(open(scaler_path)) |
| saved_window = scaler_data.get("window") |
| if saved_window is not None and saved_window != args.window: |
| raise ValueError( |
| f"Window mismatch: {scaler_path.name} was built with " |
| f"window={saved_window}, but --window={args.window} was passed. " |
| f"Re-run preprocessing.py --window {args.window}, or pass " |
| f"--window {saved_window} here to match it." |
| ) |
| n_features = X_train.shape[2] |
| model = build_model_cicids2018( |
| n_features=n_features, |
| hidden_size=args.hidden, |
| num_layers=args.layers, |
| seq_len=args.window, |
| ) |
| else: |
| scaler_path = data_dir / f"scaler_{run_id}.json" |
| scaler_data = json.load(open(scaler_path)) |
| saved_window = scaler_data.get("window") |
| if saved_window is not None and saved_window != args.window: |
| raise ValueError( |
| f"Window mismatch: {scaler_path.name} was built with " |
| f"window={saved_window}, but --window={args.window} was passed. " |
| f"Re-run preprocessing.py --window {args.window}, or pass " |
| f"--window {saved_window} here to match it." |
| ) |
| n_features = X_train.shape[2] |
| model = build_model_unsw( |
| n_features=n_features, |
| hidden_size=args.hidden, |
| num_layers=args.layers, |
| seq_len=args.window, |
| ) |
|
|
| total_params = sum(p.numel() for p in model.parameters()) |
| log.info("Model parameters: %d", total_params) |
|
|
| |
| train_losses, val_losses = train( |
| model=model, |
| X_train=X_train, |
| dataset=args.dataset, |
| run_id=run_id, |
| epochs=args.epochs, |
| batch_size=args.batch_size, |
| lr=args.lr, |
| device=device, |
| ) |
|
|
| |
| model.load_state_dict( |
| torch.load(model_dir / f"best_{run_id}.pt", map_location=device) |
| ) |
|
|
| threshold, mean_err, std_err = calculate_threshold( |
| model=model, |
| X_train=X_train, |
| dataset=args.dataset, |
| percentile=args.percentile, |
| device=device, |
| ) |
|
|
| |
| results = { |
| "dataset": args.dataset, |
| "window": args.window, |
| "threshold": threshold, |
| "mean_error": mean_err, |
| "std_error": std_err, |
| "percentile": args.percentile, |
| "epochs_trained": len(train_losses), |
| "final_loss": train_losses[-1], |
| "hidden_size": args.hidden, |
| "num_layers": args.layers, |
| "total_params": total_params, |
| } |
|
|
| |
| history = { |
| "dataset": args.dataset, |
| "window": args.window, |
| "epochs": list(range(1, len(train_losses) + 1)), |
| "train_loss": train_losses, |
| "val_loss": val_losses, |
| } |
| history_path = model_dir / f"history_{run_id}.json" |
| with open(history_path, "w") as f: |
| json.dump(history, f, indent=2) |
| log.info("History saved → %s", history_path) |
|
|
| out_path = model_dir / f"threshold_{run_id}.json" |
| with open(out_path, "w") as f: |
| json.dump(results, f, indent=2) |
|
|
| log.info("Threshold saved → %s", out_path) |
| log.info("=" * 50) |
| log.info("TRAINING COMPLETE") |
| log.info(" Best model : models/best_%s.pt", run_id) |
| log.info(" Threshold : %.6f", threshold) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|