ARotting commited on
Commit
3de20f8
·
verified ·
1 Parent(s): 272caf0

Publish 802 parameter graph compromise detector

Browse files
README.md ADDED
@@ -0,0 +1,43 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ license: apache-2.0
3
+ task_categories:
4
+ - tabular-classification
5
+ tags:
6
+ - graph-neural-network
7
+ - cybersecurity
8
+ - node-classification
9
+ - pytorch
10
+ ---
11
+
12
+ # MeshGraph GCN
13
+
14
+ MeshGraph is a graph convolutional node classifier for detecting compromised assets
15
+ inside a simulated enterprise network. Nodes carry host telemetry, while edges encode
16
+ observed communication. Compromise begins at sparse seeds and propagates stochastically
17
+ through the network.
18
+
19
+ The GCN is compared with a logistic-regression baseline that sees identical host
20
+ features but cannot use graph structure.
21
+
22
+ This is a transductive benchmark: the graph and all node features are visible during
23
+ training, while validation and test labels remain hidden.
24
+
25
+ ## Reproduce
26
+
27
+ ```powershell
28
+ uv run python projects/meshgraph-gcn/generate_data.py
29
+ uv run python projects/meshgraph-gcn/train.py
30
+ ```
31
+
32
+ ## Verified results
33
+
34
+ The generated graph contains 600 nodes, 1,689 undirected communication edges, six
35
+ subnets, and a 23.33% compromise rate. The final test contains 120 nodes:
36
+
37
+ | Model | Parameters | Accuracy | ROC-AUC | Average precision | F1 |
38
+ | --- | ---: | ---: | ---: | ---: | ---: |
39
+ | Feature-only logistic regression | 9 fitted coefficients | 78.33% | 0.8362 | 0.6274 | 0.5938 |
40
+ | Two-layer GCN | 802 | **81.67%** | **0.8564** | **0.6976** | **0.6333** |
41
+
42
+ Both thresholds were selected independently on the same 120-node validation split.
43
+ The GCN improved F1 by 3.96 points and average precision by 7.02 points.
evaluation.json ADDED
@@ -0,0 +1,45 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "model": "MeshGraph GCN",
3
+ "parameters": 802,
4
+ "nodes": 600,
5
+ "edges": 1689,
6
+ "best_epoch": 26,
7
+ "best_validation_roc_auc": 0.9118788819875777,
8
+ "gcn_threshold": 0.6789770509686359,
9
+ "gcn_test": {
10
+ "accuracy": 0.8166666666666667,
11
+ "roc_auc": 0.8563664596273292,
12
+ "average_precision": 0.6975512110914732,
13
+ "precision": 0.59375,
14
+ "recall": 0.6785714285714286,
15
+ "f1": 0.6333333333333333,
16
+ "confusion_matrix": [
17
+ [
18
+ 79,
19
+ 13
20
+ ],
21
+ [
22
+ 9,
23
+ 19
24
+ ]
25
+ ]
26
+ },
27
+ "feature_only_logistic_test": {
28
+ "accuracy": 0.7833333333333333,
29
+ "roc_auc": 0.8361801242236025,
30
+ "average_precision": 0.6273562320739494,
31
+ "precision": 0.5277777777777778,
32
+ "recall": 0.6785714285714286,
33
+ "f1": 0.59375,
34
+ "confusion_matrix": [
35
+ [
36
+ 75,
37
+ 17
38
+ ],
39
+ [
40
+ 9,
41
+ 19
42
+ ]
43
+ ]
44
+ }
45
+ }
model.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:3eec69401c9b743d6a0c72fa475710a847ba695ece1e31dc2517accd3a1492f7
3
+ size 3512
preprocessing.npz ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:87c61d8ec7ed2cd317b5becb844b767121937b1f35dae27f723155897b4d4f66
3
+ size 568
source/generate_data.py ADDED
@@ -0,0 +1,92 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ import json
4
+ from pathlib import Path
5
+
6
+ import numpy as np
7
+ from sklearn.model_selection import train_test_split
8
+
9
+ PROJECT_DIR = Path(__file__).resolve().parent
10
+ DATA_DIR = PROJECT_DIR / "data"
11
+
12
+
13
+ def build_graph(nodes: int = 600, seed: int = 2033) -> dict[str, np.ndarray]:
14
+ rng = np.random.default_rng(seed)
15
+ subnets = np.repeat(np.arange(6), nodes // 6)
16
+ adjacency = np.zeros((nodes, nodes), dtype=np.float32)
17
+ for left in range(nodes):
18
+ same_subnet = subnets == subnets[left]
19
+ probabilities = np.where(same_subnet, 0.045, 0.0025)
20
+ links = rng.random(nodes) < probabilities
21
+ links[: left + 1] = False
22
+ adjacency[left, links] = 1
23
+ adjacency = np.maximum(adjacency, adjacency.T)
24
+
25
+ compromised = np.zeros(nodes, dtype=bool)
26
+ seeds = rng.choice(nodes, size=14, replace=False)
27
+ compromised[seeds] = True
28
+ for _ in range(4):
29
+ exposure = adjacency @ compromised.astype(np.float32)
30
+ infection_probability = 1 - np.exp(-0.22 * exposure)
31
+ new_infections = (rng.random(nodes) < infection_probability) & ~compromised
32
+ compromised |= new_infections
33
+
34
+ base = rng.normal(0, 1, (nodes, 8)).astype(np.float32)
35
+ labels = compromised.astype(np.int64)
36
+ signal = labels[:, None].astype(np.float32)
37
+ features = base.copy()
38
+ features[:, 0:1] += signal * rng.normal(1.0, 0.5, (nodes, 1))
39
+ features[:, 1:2] += signal * rng.normal(0.8, 0.6, (nodes, 1))
40
+ features[:, 2:3] += signal * rng.normal(0.7, 0.6, (nodes, 1))
41
+ features[:, 3:4] += signal * rng.normal(0.5, 0.7, (nodes, 1))
42
+ features[:, 4] += subnets * 0.12
43
+
44
+ indices = np.arange(nodes)
45
+ train, remainder = train_test_split(
46
+ indices,
47
+ test_size=0.40,
48
+ stratify=labels,
49
+ random_state=seed,
50
+ )
51
+ validation, test = train_test_split(
52
+ remainder,
53
+ test_size=0.50,
54
+ stratify=labels[remainder],
55
+ random_state=seed,
56
+ )
57
+ return {
58
+ "features": features,
59
+ "adjacency": adjacency,
60
+ "labels": labels,
61
+ "subnets": subnets.astype(np.int64),
62
+ "train_indices": train,
63
+ "validation_indices": validation,
64
+ "test_indices": test,
65
+ "seed_nodes": seeds.astype(np.int64),
66
+ }
67
+
68
+
69
+ def main() -> None:
70
+ DATA_DIR.mkdir(parents=True, exist_ok=True)
71
+ graph = build_graph()
72
+ np.savez_compressed(DATA_DIR / "meshgraph.npz", **graph)
73
+ manifest = {
74
+ "nodes": len(graph["labels"]),
75
+ "edges": int(graph["adjacency"].sum() // 2),
76
+ "features": graph["features"].shape[1],
77
+ "subnets": len(np.unique(graph["subnets"])),
78
+ "compromised_rate": float(graph["labels"].mean()),
79
+ "train_nodes": len(graph["train_indices"]),
80
+ "validation_nodes": len(graph["validation_indices"]),
81
+ "test_nodes": len(graph["test_indices"]),
82
+ "path": "meshgraph.npz",
83
+ }
84
+ (DATA_DIR / "manifest.json").write_text(
85
+ json.dumps(manifest, indent=2),
86
+ encoding="utf-8",
87
+ )
88
+ print(json.dumps(manifest, indent=2))
89
+
90
+
91
+ if __name__ == "__main__":
92
+ main()
source/model.py ADDED
@@ -0,0 +1,36 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ import torch
4
+ from torch import nn
5
+ from torch.nn import functional as F
6
+
7
+
8
+ class MeshGraphGCN(nn.Module):
9
+ def __init__(self, features: int = 8) -> None:
10
+ super().__init__()
11
+ self.input = nn.Linear(features, 32, bias=False)
12
+ self.hidden = nn.Linear(32, 16, bias=False)
13
+ self.output = nn.Linear(16, 2)
14
+
15
+ def forward(
16
+ self,
17
+ features: torch.Tensor,
18
+ normalized_adjacency: torch.Tensor,
19
+ ) -> torch.Tensor:
20
+ hidden = normalized_adjacency @ features
21
+ hidden = F.gelu(self.input(hidden))
22
+ hidden = F.dropout(hidden, p=0.15, training=self.training)
23
+ hidden = normalized_adjacency @ hidden
24
+ hidden = F.gelu(self.hidden(hidden))
25
+ return self.output(hidden)
26
+
27
+
28
+ def normalize_adjacency(adjacency: torch.Tensor) -> torch.Tensor:
29
+ with_self_loops = adjacency + torch.eye(len(adjacency))
30
+ degree = with_self_loops.sum(dim=1).clamp(min=1)
31
+ inverse_sqrt = degree.pow(-0.5)
32
+ return inverse_sqrt[:, None] * with_self_loops * inverse_sqrt[None, :]
33
+
34
+
35
+ def parameter_count(model: nn.Module) -> int:
36
+ return sum(parameter.numel() for parameter in model.parameters())
source/train.py ADDED
@@ -0,0 +1,191 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ import json
4
+ import random
5
+ from pathlib import Path
6
+
7
+ import numpy as np
8
+ import torch
9
+ import trackio
10
+ from model import MeshGraphGCN, normalize_adjacency, parameter_count
11
+ from safetensors.torch import save_file
12
+ from sklearn.linear_model import LogisticRegression
13
+ from sklearn.metrics import (
14
+ accuracy_score,
15
+ average_precision_score,
16
+ confusion_matrix,
17
+ f1_score,
18
+ precision_score,
19
+ recall_score,
20
+ roc_auc_score,
21
+ )
22
+ from torch.nn import functional as F
23
+
24
+ PROJECT_DIR = Path(__file__).resolve().parent
25
+ DATA_DIR = PROJECT_DIR / "data"
26
+ ARTIFACT_DIR = PROJECT_DIR / "artifacts" / "meshgraph-gcn"
27
+
28
+
29
+ def seed_everything(seed: int) -> None:
30
+ random.seed(seed)
31
+ np.random.seed(seed)
32
+ torch.manual_seed(seed)
33
+
34
+
35
+ def metrics(labels: np.ndarray, scores: np.ndarray, threshold: float = 0.5) -> dict:
36
+ predictions = scores >= threshold
37
+ return {
38
+ "accuracy": float(accuracy_score(labels, predictions)),
39
+ "roc_auc": float(roc_auc_score(labels, scores)),
40
+ "average_precision": float(average_precision_score(labels, scores)),
41
+ "precision": float(precision_score(labels, predictions, zero_division=0)),
42
+ "recall": float(recall_score(labels, predictions, zero_division=0)),
43
+ "f1": float(f1_score(labels, predictions, zero_division=0)),
44
+ "confusion_matrix": confusion_matrix(labels, predictions).tolist(),
45
+ }
46
+
47
+
48
+ def best_threshold(labels: np.ndarray, scores: np.ndarray) -> float:
49
+ candidates = np.quantile(scores, np.linspace(0.05, 0.95, 300))
50
+ return float(
51
+ max(
52
+ candidates,
53
+ key=lambda threshold: f1_score(
54
+ labels,
55
+ scores >= threshold,
56
+ zero_division=0,
57
+ ),
58
+ )
59
+ )
60
+
61
+
62
+ def main() -> None:
63
+ seed_everything(2033)
64
+ graph = np.load(DATA_DIR / "meshgraph.npz")
65
+ features = graph["features"].astype(np.float32)
66
+ labels = graph["labels"].astype(np.int64)
67
+ train_indices = graph["train_indices"]
68
+ validation_indices = graph["validation_indices"]
69
+ test_indices = graph["test_indices"]
70
+ mean = features[train_indices].mean(axis=0)
71
+ scale = np.maximum(features[train_indices].std(axis=0), 1e-5)
72
+ features = (features - mean) / scale
73
+ feature_tensor = torch.from_numpy(features)
74
+ label_tensor = torch.from_numpy(labels)
75
+ adjacency = torch.from_numpy(graph["adjacency"].astype(np.float32))
76
+ normalized = normalize_adjacency(adjacency)
77
+
78
+ baseline = LogisticRegression(
79
+ class_weight="balanced",
80
+ max_iter=2000,
81
+ random_state=2033,
82
+ )
83
+ baseline.fit(features[train_indices], labels[train_indices])
84
+ baseline_validation = baseline.predict_proba(features[validation_indices])[:, 1]
85
+ baseline_threshold = best_threshold(
86
+ labels[validation_indices],
87
+ baseline_validation,
88
+ )
89
+ baseline_test = baseline.predict_proba(features[test_indices])[:, 1]
90
+
91
+ model = MeshGraphGCN(features=features.shape[1])
92
+ positive_weight = (labels[train_indices] == 0).sum() / max(
93
+ 1, (labels[train_indices] == 1).sum()
94
+ )
95
+ class_weights = torch.tensor([1.0, positive_weight], dtype=torch.float32)
96
+ optimizer = torch.optim.AdamW(model.parameters(), lr=0.01, weight_decay=0.002)
97
+ scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=350)
98
+ best_validation_auc = -1.0
99
+ best_epoch = 0
100
+ best_state = None
101
+ trackio.init(
102
+ project="meshgraph-gcn",
103
+ name="two-layer-gcn-v1",
104
+ config={
105
+ "parameters": parameter_count(model),
106
+ "nodes": len(labels),
107
+ "edges": int(adjacency.sum().item() // 2),
108
+ "train_labels": len(train_indices),
109
+ "transductive": True,
110
+ },
111
+ )
112
+ for epoch in range(1, 351):
113
+ model.train()
114
+ logits = model(feature_tensor, normalized)
115
+ loss = F.cross_entropy(
116
+ logits[train_indices],
117
+ label_tensor[train_indices],
118
+ weight=class_weights,
119
+ )
120
+ optimizer.zero_grad(set_to_none=True)
121
+ loss.backward()
122
+ optimizer.step()
123
+ scheduler.step()
124
+ model.eval()
125
+ with torch.no_grad():
126
+ probabilities = model(feature_tensor, normalized).softmax(dim=1)[:, 1]
127
+ validation_auc = roc_auc_score(
128
+ labels[validation_indices],
129
+ probabilities[validation_indices].numpy(),
130
+ )
131
+ if validation_auc > best_validation_auc:
132
+ best_validation_auc = validation_auc
133
+ best_epoch = epoch
134
+ best_state = {
135
+ key: value.detach().cpu().clone()
136
+ for key, value in model.state_dict().items()
137
+ }
138
+ if epoch == 1 or epoch % 10 == 0:
139
+ trackio.log(
140
+ {
141
+ "epoch": epoch,
142
+ "train_loss": float(loss.detach()),
143
+ "validation_roc_auc": validation_auc,
144
+ "learning_rate": scheduler.get_last_lr()[0],
145
+ }
146
+ )
147
+ trackio.finish()
148
+ assert best_state is not None
149
+ model.load_state_dict(best_state)
150
+ model.eval()
151
+ with torch.no_grad():
152
+ gcn_scores = model(feature_tensor, normalized).softmax(dim=1)[:, 1].numpy()
153
+ gcn_threshold = best_threshold(
154
+ labels[validation_indices],
155
+ gcn_scores[validation_indices],
156
+ )
157
+ results = {
158
+ "model": "MeshGraph GCN",
159
+ "parameters": parameter_count(model),
160
+ "nodes": len(labels),
161
+ "edges": int(adjacency.sum().item() // 2),
162
+ "best_epoch": best_epoch,
163
+ "best_validation_roc_auc": best_validation_auc,
164
+ "gcn_threshold": gcn_threshold,
165
+ "gcn_test": metrics(
166
+ labels[test_indices],
167
+ gcn_scores[test_indices],
168
+ gcn_threshold,
169
+ ),
170
+ "feature_only_logistic_test": metrics(
171
+ labels[test_indices],
172
+ baseline_test,
173
+ baseline_threshold,
174
+ ),
175
+ }
176
+ ARTIFACT_DIR.mkdir(parents=True, exist_ok=True)
177
+ save_file(model.state_dict(), ARTIFACT_DIR / "model.safetensors")
178
+ np.savez(
179
+ ARTIFACT_DIR / "preprocessing.npz",
180
+ mean=mean,
181
+ scale=scale,
182
+ )
183
+ (ARTIFACT_DIR / "evaluation.json").write_text(
184
+ json.dumps(results, indent=2),
185
+ encoding="utf-8",
186
+ )
187
+ print(json.dumps(results, indent=2))
188
+
189
+
190
+ if __name__ == "__main__":
191
+ main()