atharvawarade9807's picture
Upload 18 files
c6e6f10 verified
Raw
History Blame Contribute Delete
2.63 kB
import json
import torch
import torch.nn as nn
import torch.optim as optim
from models.gnn import PhishingGNN_Model
from pipeline.graph_engine import TopologicalGraphEngine
from config import *
def run_training_pipeline():
print("[*] Launching Production Defender V5 Training Run...")
# 1. Try to load real data, fallback to mock if missing
try:
with open("network_telemetry.json", "r") as f:
dataset = json.load(f)
print(f"[+] Loaded {len(dataset)} real telemetry logs.")
except FileNotFoundError:
print("[-] network_telemetry.json not found. Falling back to simulated context.")
dataset = [
{'ip': '185.220.101.5', 'domain': 'shadow-dns-bypass.net', 'asn': 44050, 'is_malicious': 1.0},
{'ip': '193.56.28.14', 'asn': 57099, 'is_malicious': 1.0},
{'ip': '8.8.8.8', 'domain': 'dns.google', 'asn': 15169, 'is_malicious': 0.0}
]
# 2. Extract logs and structure target threat labels
engine = TopologicalGraphEngine()
x_dict, edge_index_dict = engine.extract_and_build(dataset)
labels_list = []
for ip_str in engine.ip_map.keys():
matching_logs = [log for log in dataset if log.get('ip') == ip_str]
label = matching_logs[0].get('is_malicious', 0.0) if matching_logs else 0.0
labels_list.append([label])
labels = torch.tensor(labels_list, dtype=torch.float32)
# 3. Model setup
in_channels_dict = {'ip': 16, 'domain': 32, 'asn': 8, 'cert': 16}
model = PhishingGNN_Model(
metadata=GRAPH_METADATA,
in_channels_dict=in_channels_dict,
hidden_channels=HIDDEN_CHANNELS,
num_heads=NUM_HEADS,
num_layers=NUM_LAYERS,
dropout_rate=DROPOUT_RATE
)
optimizer = optim.AdamW(model.parameters(), lr=0.0005, weight_decay=1e-3)
criterion = nn.BCELoss()
# 4. Optimization Loop
model.train()
for epoch in range(100):
optimizer.zero_grad()
predictions = model(x_dict, edge_index_dict)
# Only calculate loss on the nodes we have labels for
valid_preds = predictions[:len(labels)]
loss = criterion(valid_preds, labels)
loss.backward()
optimizer.step()
if (epoch + 1) % 20 == 0:
print(f"Epoch {epoch+1:03d}/100 | Topological Loss: {loss.item():.5f}")
model.safe_save(MODEL_SAVE_PATH)
print(f"[+] Operational weights successfully frozen at: {MODEL_SAVE_PATH}")
if __name__ == "__main__":
run_training_pipeline()