| """ |
| evaluate_balanced.py |
| ===================== |
| Evaluates all three training variants on the SAME balanced test set |
| (X_test_balanced_{run_id}.npy / y_test_balanced_{run_id}.npy produced |
| by balance_data.py), so the comparison is apples-to-apples: |
| |
| 1. normal_only -- your existing canonical model |
| (models/best_{run_id}.pt), re-evaluated |
| here on the NEW balanced test set (its |
| own Chapter 4 numbers, from the ORIGINAL |
| test set, are untouched and reported |
| separately -- this is an extra data |
| point, not a replacement). |
| 2. balanced_recon -- models/best_{run_id}_balanced_recon.pt |
| 3. balanced_supervised -- models/best_{run_id}_balanced_supervised.pt |
| |
| Reconstruction-error models (1 and 2) still need a threshold. Rather |
| than re-deriving a fresh percentile threshold on the new balanced |
| training data (which would conflate "did balancing help" with "did |
| recalibrating the threshold help"), this script calibrates the |
| threshold for each reconstruction-based model the same way as |
| train.py: 95th percentile of that model's OWN reconstruction error |
| on ITS OWN normal training sessions. For normal_only that's the |
| original X_train; for balanced_recon that's the normal subset of |
| X_train_balanced. This keeps the calibration methodology identical |
| across variants -- only the training data composition differs. |
| |
| The supervised variant doesn't need a threshold at all -- it outputs |
| a probability directly from its classifier head (threshold = 0.5). |
| |
| Usage |
| ----- |
| python evaluate_balanced.py --dataset csic2010 --window 5 |
| |
| Author : K.A.D.S.D. Kandanaarachchi (2020/ICT/19) |
| Project: Detecting Anomalous REST API Traffic -- IT4216 (balanced- |
| training ablation) |
| """ |
|
|
| import argparse |
| import json |
| import logging |
| from pathlib import Path |
|
|
| import numpy as np |
| import torch |
| from sklearn.metrics import confusion_matrix, f1_score, precision_score, recall_score |
| from torch.utils.data import DataLoader, TensorDataset |
|
|
| from model import build_model_cicids2018, build_model_csic2010, build_model_unsw |
| from train_balanced import LSTMAutoencoderWithHead, _to_tensor |
|
|
| logging.basicConfig( |
| level=logging.INFO, |
| format="%(asctime)s %(levelname)s %(message)s", |
| datefmt="%H:%M:%S", |
| ) |
| log = logging.getLogger(__name__) |
|
|
|
|
| def compute_metrics(y_true, y_pred, name): |
| prec = precision_score(y_true, y_pred, zero_division=0) |
| rec = recall_score(y_true, y_pred, zero_division=0) |
| f1 = f1_score(y_true, y_pred, zero_division=0) |
| cm = confusion_matrix(y_true, y_pred) |
| tn, fp, fn, tp = cm.ravel() |
| fpr = fp / (fp + tn) if (fp + tn) > 0 else 0.0 |
| log.info( |
| "%-20s Prec=%.4f Rec=%.4f F1=%.4f FPR=%.4f TP=%d FP=%d TN=%d FN=%d", |
| name, prec, rec, f1, fpr, tp, fp, tn, fn, |
| ) |
| return { |
| "model": name, "precision": round(prec, 4), "recall": round(rec, 4), |
| "f1": round(f1, 4), "fpr": round(fpr, 4), |
| "tp": int(tp), "fp": int(fp), "tn": int(tn), "fn": int(fn), |
| } |
|
|
|
|
| def reconstruction_errors(model, X, dataset, device, batch_size=512): |
| tensor = _to_tensor(X, dataset) |
| loader = DataLoader(TensorDataset(tensor), batch_size=batch_size) |
| errors = [] |
| with torch.no_grad(): |
| for (batch,) in loader: |
| batch = batch.to(device) |
| errors.extend(model.reconstruction_error(batch).cpu().numpy()) |
| return np.array(errors) |
|
|
|
|
| def calibrate_threshold(model, X_normal_train, dataset, device, percentile=95.0): |
| errors = reconstruction_errors(model, X_normal_train, dataset, device) |
| return float(np.percentile(errors, percentile)) |
|
|
|
|
| def build_base_model(dataset, window, data_dir, X_shape): |
| if dataset == "csic2010": |
| vocab_data = json.load(open(data_dir / f"vocab_{dataset}_w{window}.json")) |
| vocab = vocab_data.get("vocab", vocab_data) |
| return build_model_csic2010(vocab_size=len(vocab), seq_len=window) |
| elif dataset == "cicids2018": |
| return build_model_cicids2018(n_features=X_shape[2], seq_len=window) |
| else: |
| return build_model_unsw(n_features=X_shape[2], seq_len=window) |
|
|
|
|
| def main(): |
| parser = argparse.ArgumentParser() |
| parser.add_argument("--dataset", required=True, choices=["csic2010", "cicids2018", "unsw"]) |
| parser.add_argument("--window", type=int, default=5) |
| args = parser.parse_args() |
|
|
| device = "cuda" if torch.cuda.is_available() else "cpu" |
| data_dir = Path("data/processed") |
| model_dir = Path("models") |
| res_dir = Path("results") |
| res_dir.mkdir(exist_ok=True) |
| run_id = f"{args.dataset}_w{args.window}" |
|
|
| X_test_bal = np.load(data_dir / f"X_test_balanced_{run_id}.npy") |
| y_test_bal = np.load(data_dir / f"y_test_balanced_{run_id}.npy") |
| X_train_bal = np.load(data_dir / f"X_train_balanced_{run_id}.npy") |
| y_train_bal = np.load(data_dir / f"y_train_balanced_{run_id}.npy") |
| X_train_normal_orig = np.load(data_dir / f"X_train_{run_id}.npy") |
|
|
| log.info( |
| "Balanced test set: %d normal / %d attack", |
| int((y_test_bal == 0).sum()), int((y_test_bal == 1).sum()), |
| ) |
|
|
| results = [] |
|
|
| |
| |
| |
| |
| normal_only_path = model_dir / f"best_{run_id}.pt" |
| if normal_only_path.exists(): |
| base = build_base_model(args.dataset, args.window, data_dir, X_test_bal.shape) |
| base.load_state_dict(torch.load(normal_only_path, map_location=device)) |
| base = base.to(device) |
| base.eval() |
| thresh = calibrate_threshold(base, X_train_normal_orig, args.dataset, device) |
| errors = reconstruction_errors(base, X_test_bal, args.dataset, device) |
| y_pred = (errors > thresh).astype(int) |
| results.append(compute_metrics(y_test_bal, y_pred, "normal_only (on balanced test)")) |
| else: |
| log.warning(" %s not found -- skipping normal_only comparison", normal_only_path) |
|
|
| |
| |
| recon_path = model_dir / f"best_{run_id}_balanced_recon.pt" |
| if recon_path.exists(): |
| base2 = build_base_model(args.dataset, args.window, data_dir, X_test_bal.shape) |
| base2.load_state_dict(torch.load(recon_path, map_location=device)) |
| base2 = base2.to(device) |
| base2.eval() |
| X_train_bal_normal_only = X_train_bal[y_train_bal == 0] |
| thresh2 = calibrate_threshold(base2, X_train_bal_normal_only, args.dataset, device) |
| errors2 = reconstruction_errors(base2, X_test_bal, args.dataset, device) |
| y_pred2 = (errors2 > thresh2).astype(int) |
| results.append(compute_metrics(y_test_bal, y_pred2, "balanced_recon")) |
| else: |
| log.warning(" %s not found -- run train_balanced.py --mode balanced_recon first", recon_path) |
|
|
| |
| |
| sup_path = model_dir / f"best_{run_id}_balanced_supervised.pt" |
| if sup_path.exists(): |
| base3 = build_base_model(args.dataset, args.window, data_dir, X_test_bal.shape) |
| model3 = LSTMAutoencoderWithHead(base3, hidden_size=64) |
| model3.load_state_dict(torch.load(sup_path, map_location=device)) |
| model3 = model3.to(device) |
| model3.eval() |
|
|
| tensor = _to_tensor(X_test_bal, args.dataset) |
| loader = DataLoader(TensorDataset(tensor), batch_size=512) |
| probs = [] |
| with torch.no_grad(): |
| for (batch,) in loader: |
| batch = batch.to(device) |
| probs.extend(model3.predict_proba(batch).cpu().numpy()) |
| probs = np.array(probs) |
| y_pred3 = (probs > 0.5).astype(int) |
| results.append(compute_metrics(y_test_bal, y_pred3, "balanced_supervised")) |
| else: |
| log.warning(" %s not found -- run train_balanced.py --mode balanced_supervised first", sup_path) |
|
|
| if not results: |
| log.error("No trained balanced-variant models found. Run train_balanced.py first.") |
| return |
|
|
| out = { |
| "dataset": args.dataset, |
| "window": args.window, |
| "test_set": "balanced (see balance_data.py summary for exact composition)", |
| "test_normal": int((y_test_bal == 0).sum()), |
| "test_attack": int((y_test_bal == 1).sum()), |
| "results": results, |
| } |
| out_path = res_dir / f"evaluation_balanced_{run_id}.json" |
| with open(out_path, "w") as f: |
| json.dump(out, f, indent=2) |
| log.info("Saved -> %s", out_path) |
|
|
| log.info("") |
| log.info("=" * 65) |
| log.info("BALANCED-TRAINING COMPARISON -- %s", args.dataset.upper()) |
| log.info("=" * 65) |
| log.info("%-32s %6s %6s %6s %6s", "Model", "Prec", "Rec", "F1", "FPR") |
| for r in results: |
| log.info("%-32s %6.4f %6.4f %6.4f %6.4f", r["model"], r["precision"], r["recall"], r["f1"], r["fpr"]) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|