""" 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 90% train / 10% validation split = int(len(X_train) * 0.9) X_tr = X_train[:split] X_val = X_train[split:] # Convert to tensors 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): # Training model.train() epoch_loss = 0.0 for (batch,) in tr_loader: batch = batch.to(device) optimizer.zero_grad() recon = model(batch) # For embedding mode, compare against embedded input if dataset == "csic2010": target = model.embedding(batch).detach() else: target = batch loss = criterion(recon, target) loss.backward() # Gradient clipping — prevents exploding gradients in LSTMs nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) optimizer.step() epoch_loss += loss.item() * len(batch) epoch_loss /= len(X_tr) # Validation 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 ) # Early stopping if val_loss < best_val - 1e-6: best_val = val_loss patience_ctr = 0 # Save best weights 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) # Load data — filenames are window-suffixed by preprocessing.py X_train = np.load(data_dir / f"X_train_{run_id}.npy") log.info("X_train shape: %s", X_train.shape) # Build model — also cross-check the window baked into # preprocessing's metadata against --window, so a mismatched flag # errors out loudly instead of silently building the wrong seq_len. 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) # back-compat with old bare-dict format 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 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, ) # Load best weights and calculate threshold 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, ) # Save results 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, } # Save per-epoch history for visualisation 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()