DilanjanaSDK's picture
Upload folder using huggingface_hub
c76fc30 verified
Raw
History Blame Contribute Delete
11.8 kB
"""
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()