LogSentry-LM / detect_anomalies.py
sankath's picture
Upload detect_anomalies.py with huggingface_hub
d0d5435 verified
Raw
History Blame Contribute Delete
3.01 kB
# file: detect_anomalies.py
# Standalone predictor: load the trained model + saved threshold and flag lines
# the model finds surprising. Run after train_and_generate.py has saved a checkpoint.
# python detect_anomalies.py "Failed password for root from 1.2.3.4 port 22 ssh2"
# python detect_anomalies.py --file some_logs.txt
# python detect_anomalies.py # built-in demo lines
import os
import sys
import argparse
import torch
from dataset import build_or_load_tokenizer
from train_and_generate import LogSentryLM, score_line
MODEL_FILE = "model/logsentry_lm.pt"
TOKENIZER_FILE = "tokenizer/log_tokenizer.json"
LOG_FILE = "data/raw_logs.txt"
def pick_device():
if torch.cuda.is_available():
return "cuda"
if torch.backends.mps.is_available():
return "mps"
return "cpu"
def load_model(device):
if not os.path.exists(MODEL_FILE):
sys.exit(f"[ERROR] Model checkpoint '{MODEL_FILE}' not found. "
f"Run `python train_and_generate.py` first to train and save it.")
ckpt = torch.load(MODEL_FILE, map_location=device)
# Rebuild the exact architecture from the saved config, then load weights
model = LogSentryLM(vocab_size=ckpt["vocab_size"], max_seq_len=ckpt["max_seq_len"]).to(device)
model.load_state_dict(ckpt["model_state"])
model.eval()
tokenizer = build_or_load_tokenizer(LOG_FILE, TOKENIZER_FILE)
return model, tokenizer, ckpt
def predict(model, tokenizer, lines, ckpt, device):
threshold = ckpt["threshold"]
max_seq_len = ckpt["max_seq_len"]
print(f"[DETECT] Anomaly threshold (loss > {threshold:.3f} => ANOMALY)\n")
for line in lines:
result = score_line(model, tokenizer, line, device, max_seq_len=max_seq_len)
if result is None:
print(f" [skipped: too short] {line}")
continue
loss, ppl = result
flag = "ANOMALY" if loss > threshold else "normal "
print(f" [{flag}] loss={loss:.3f} ppl={ppl:9.1f} | {line[:100]}")
if __name__ == "__main__":
parser = argparse.ArgumentParser(description="Flag anomalous log lines with LogSentry-LM.")
parser.add_argument("lines", nargs="*", help="One or more log lines to score.")
parser.add_argument("--file", help="Path to a log file; scores each non-empty line.")
args = parser.parse_args()
device = pick_device()
model, tokenizer, ckpt = load_model(device)
if args.file:
with open(args.file, "r") as f:
lines = [ln.strip() for ln in f if ln.strip()]
elif args.lines:
lines = args.lines
else:
# demo: two genuine lines + one clearly off-distribution line
lines = [
"Failed password for root from 112.95.230.3 port 49204 ssh2",
"[Sun Dec 04 04:47:44 2005] [notice] workerEnv.init() ok /etc/httpd/conf/workers2.properties",
"kjshdf!! TOTALLY RANDOM gibberish $$$ 999 not-a-real-log ~~~~",
]
predict(model, tokenizer, lines, ckpt, device)