@vk.amogh/trace 2.1.0
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- package/README.md +208 -0
- package/bin/trace.js +112 -0
- package/package.json +46 -0
- package/pyproject.toml +39 -0
- package/src/trace_engine/__init__.py +8 -0
- package/src/trace_engine/__pycache__/__init__.cpython-311.pyc +0 -0
- package/src/trace_engine/__pycache__/cli.cpython-311.pyc +0 -0
- package/src/trace_engine/__pycache__/doctor.cpython-311.pyc +0 -0
- package/src/trace_engine/__pycache__/interactive.cpython-311.pyc +0 -0
- package/src/trace_engine/__pycache__/verify.cpython-311.pyc +0 -0
- package/src/trace_engine/ai/__init__.py +7 -0
- package/src/trace_engine/ai/__pycache__/__init__.cpython-311.pyc +0 -0
- package/src/trace_engine/ai/__pycache__/base.cpython-311.pyc +0 -0
- package/src/trace_engine/ai/__pycache__/ollama.cpython-311.pyc +0 -0
- package/src/trace_engine/ai/__pycache__/planner.cpython-311.pyc +0 -0
- package/src/trace_engine/ai/base.py +23 -0
- package/src/trace_engine/ai/ollama.py +50 -0
- package/src/trace_engine/ai/planner.py +40 -0
- package/src/trace_engine/apm/__init__.py +18 -0
- package/src/trace_engine/apm/__pycache__/__init__.cpython-311.pyc +0 -0
- package/src/trace_engine/apm/__pycache__/builder.cpython-311.pyc +0 -0
- package/src/trace_engine/apm/__pycache__/edges.cpython-311.pyc +0 -0
- package/src/trace_engine/apm/__pycache__/model.cpython-311.pyc +0 -0
- package/src/trace_engine/apm/__pycache__/nodes.cpython-311.pyc +0 -0
- package/src/trace_engine/apm/__pycache__/serialization.cpython-311.pyc +0 -0
- package/src/trace_engine/apm/builder.py +208 -0
- package/src/trace_engine/apm/edges.py +25 -0
- package/src/trace_engine/apm/model.py +107 -0
- package/src/trace_engine/apm/nodes.py +27 -0
- package/src/trace_engine/apm/serialization.py +105 -0
- package/src/trace_engine/benchmark/__init__.py +5 -0
- package/src/trace_engine/benchmark/__pycache__/__init__.cpython-311.pyc +0 -0
- package/src/trace_engine/benchmark/__pycache__/owasp.cpython-311.pyc +0 -0
- package/src/trace_engine/benchmark/owasp.py +183 -0
- package/src/trace_engine/cli.py +1184 -0
- package/src/trace_engine/config/__init__.py +35 -0
- package/src/trace_engine/config/__pycache__/__init__.cpython-311.pyc +0 -0
- package/src/trace_engine/config/__pycache__/defaults.cpython-311.pyc +0 -0
- package/src/trace_engine/config/__pycache__/loader.cpython-311.pyc +0 -0
- package/src/trace_engine/config/__pycache__/settings.cpython-311.pyc +0 -0
- package/src/trace_engine/config/defaults.py +48 -0
- package/src/trace_engine/config/loader.py +64 -0
- package/src/trace_engine/config/settings.py +72 -0
- package/src/trace_engine/doctor.py +250 -0
- package/src/trace_engine/findings/__init__.py +15 -0
- package/src/trace_engine/findings/__pycache__/__init__.cpython-311.pyc +0 -0
- package/src/trace_engine/findings/__pycache__/correlate.cpython-311.pyc +0 -0
- package/src/trace_engine/findings/__pycache__/model.cpython-311.pyc +0 -0
- package/src/trace_engine/findings/__pycache__/recommendations.cpython-311.pyc +0 -0
- package/src/trace_engine/findings/__pycache__/store.cpython-311.pyc +0 -0
- package/src/trace_engine/findings/correlate.py +103 -0
- package/src/trace_engine/findings/model.py +40 -0
- package/src/trace_engine/findings/recommendations.py +35 -0
- package/src/trace_engine/findings/store.py +39 -0
- package/src/trace_engine/framework/__init__.py +58 -0
- package/src/trace_engine/framework/__pycache__/__init__.cpython-311.pyc +0 -0
- package/src/trace_engine/framework/__pycache__/base.cpython-311.pyc +0 -0
- package/src/trace_engine/framework/__pycache__/csharp.cpython-311.pyc +0 -0
- package/src/trace_engine/framework/__pycache__/dart.cpython-311.pyc +0 -0
- package/src/trace_engine/framework/__pycache__/django.cpython-311.pyc +0 -0
- package/src/trace_engine/framework/__pycache__/express.cpython-311.pyc +0 -0
- package/src/trace_engine/framework/__pycache__/fastapi.cpython-311.pyc +0 -0
- package/src/trace_engine/framework/__pycache__/flask.cpython-311.pyc +0 -0
- package/src/trace_engine/framework/__pycache__/go.cpython-311.pyc +0 -0
- package/src/trace_engine/framework/__pycache__/nextjs.cpython-311.pyc +0 -0
- package/src/trace_engine/framework/__pycache__/php.cpython-311.pyc +0 -0
- package/src/trace_engine/framework/__pycache__/react_router.cpython-311.pyc +0 -0
- package/src/trace_engine/framework/__pycache__/ruby.cpython-311.pyc +0 -0
- package/src/trace_engine/framework/__pycache__/rust.cpython-311.pyc +0 -0
- package/src/trace_engine/framework/__pycache__/springboot.cpython-311.pyc +0 -0
- package/src/trace_engine/framework/base.py +49 -0
- package/src/trace_engine/framework/csharp.py +111 -0
- package/src/trace_engine/framework/dart.py +152 -0
- package/src/trace_engine/framework/django.py +188 -0
- package/src/trace_engine/framework/express.py +82 -0
- package/src/trace_engine/framework/fastapi.py +135 -0
- package/src/trace_engine/framework/flask.py +108 -0
- package/src/trace_engine/framework/go.py +94 -0
- package/src/trace_engine/framework/nextjs.py +200 -0
- package/src/trace_engine/framework/php.py +98 -0
- package/src/trace_engine/framework/react_router.py +331 -0
- package/src/trace_engine/framework/ruby.py +69 -0
- package/src/trace_engine/framework/rust.py +101 -0
- package/src/trace_engine/framework/springboot.py +145 -0
- package/src/trace_engine/harness/__init__.py +19 -0
- package/src/trace_engine/harness/__pycache__/__init__.cpython-311.pyc +0 -0
- package/src/trace_engine/harness/__pycache__/benchmark.cpython-311.pyc +0 -0
- package/src/trace_engine/harness/__pycache__/context.cpython-311.pyc +0 -0
- package/src/trace_engine/harness/__pycache__/engine.cpython-311.pyc +0 -0
- package/src/trace_engine/harness/__pycache__/patcher.cpython-311.pyc +0 -0
- package/src/trace_engine/harness/__pycache__/remediators.cpython-311.pyc +0 -0
- package/src/trace_engine/harness/benchmark.py +55 -0
- package/src/trace_engine/harness/context.py +49 -0
- package/src/trace_engine/harness/engine.py +187 -0
- package/src/trace_engine/harness/patcher.py +143 -0
- package/src/trace_engine/harness/remediators.py +307 -0
- package/src/trace_engine/ingest/__init__.py +15 -0
- package/src/trace_engine/ingest/__pycache__/__init__.cpython-311.pyc +0 -0
- package/src/trace_engine/ingest/__pycache__/files.cpython-311.pyc +0 -0
- package/src/trace_engine/ingest/__pycache__/hashing.cpython-311.pyc +0 -0
- package/src/trace_engine/ingest/__pycache__/ignore.cpython-311.pyc +0 -0
- package/src/trace_engine/ingest/__pycache__/repository.cpython-311.pyc +0 -0
- package/src/trace_engine/ingest/files.py +60 -0
- package/src/trace_engine/ingest/hashing.py +21 -0
- package/src/trace_engine/ingest/ignore.py +108 -0
- package/src/trace_engine/ingest/repository.py +50 -0
- package/src/trace_engine/intelligence/__init__.py +15 -0
- package/src/trace_engine/intelligence/__pycache__/__init__.cpython-311.pyc +0 -0
- package/src/trace_engine/intelligence/__pycache__/evaluation.cpython-311.pyc +0 -0
- package/src/trace_engine/intelligence/__pycache__/orchestrator.cpython-311.pyc +0 -0
- package/src/trace_engine/intelligence/__pycache__/tracebench.cpython-311.pyc +0 -0
- package/src/trace_engine/intelligence/evaluation.py +828 -0
- package/src/trace_engine/intelligence/laya/__init__.py +19 -0
- package/src/trace_engine/intelligence/laya/__pycache__/__init__.cpython-311.pyc +0 -0
- package/src/trace_engine/intelligence/laya/__pycache__/prompts.cpython-311.pyc +0 -0
- package/src/trace_engine/intelligence/laya/__pycache__/router.cpython-311.pyc +0 -0
- package/src/trace_engine/intelligence/laya/__pycache__/schemas.cpython-311.pyc +0 -0
- package/src/trace_engine/intelligence/laya/__pycache__/telemetry.cpython-311.pyc +0 -0
- package/src/trace_engine/intelligence/laya/__pycache__/thresholds.cpython-311.pyc +0 -0
- package/src/trace_engine/intelligence/laya/prompts.py +67 -0
- package/src/trace_engine/intelligence/laya/router.py +352 -0
- package/src/trace_engine/intelligence/laya/schemas.py +56 -0
- package/src/trace_engine/intelligence/laya/telemetry.py +48 -0
- package/src/trace_engine/intelligence/laya/thresholds.py +12 -0
- package/src/trace_engine/intelligence/orchestrator.py +130 -0
- package/src/trace_engine/intelligence/securebert/__init__.py +6 -0
- package/src/trace_engine/intelligence/securebert/__pycache__/__init__.cpython-311.pyc +0 -0
- package/src/trace_engine/intelligence/securebert/__pycache__/cache.cpython-311.pyc +0 -0
- package/src/trace_engine/intelligence/securebert/__pycache__/classifier.cpython-311.pyc +0 -0
- package/src/trace_engine/intelligence/securebert/cache.py +37 -0
- package/src/trace_engine/intelligence/securebert/classifier.py +240 -0
- package/src/trace_engine/intelligence/tracebench.py +61 -0
- package/src/trace_engine/intelligence/training/__init__.py +21 -0
- package/src/trace_engine/intelligence/training/__pycache__/__init__.cpython-311.pyc +0 -0
- package/src/trace_engine/intelligence/training/__pycache__/dataset.cpython-311.pyc +0 -0
- package/src/trace_engine/intelligence/training/__pycache__/dataset_importers.cpython-311.pyc +0 -0
- package/src/trace_engine/intelligence/training/__pycache__/laya_trainer.cpython-311.pyc +0 -0
- package/src/trace_engine/intelligence/training/__pycache__/lora_system2.cpython-311.pyc +0 -0
- package/src/trace_engine/intelligence/training/__pycache__/morefixes_pipeline.cpython-311.pyc +0 -0
- package/src/trace_engine/intelligence/training/__pycache__/slicer.cpython-311.pyc +0 -0
- package/src/trace_engine/intelligence/training/__pycache__/train_all.cpython-311.pyc +0 -0
- package/src/trace_engine/intelligence/training/__pycache__/trainer.cpython-311.pyc +0 -0
- package/src/trace_engine/intelligence/training/dataset.py +320 -0
- package/src/trace_engine/intelligence/training/dataset_importers.py +349 -0
- package/src/trace_engine/intelligence/training/laya_trainer.py +678 -0
- package/src/trace_engine/intelligence/training/lora_system2.py +162 -0
- package/src/trace_engine/intelligence/training/morefixes_pipeline.py +463 -0
- package/src/trace_engine/intelligence/training/slicer.py +129 -0
- package/src/trace_engine/intelligence/training/train_all.py +1009 -0
- package/src/trace_engine/intelligence/training/trainer.py +321 -0
- package/src/trace_engine/interactive.py +623 -0
- package/src/trace_engine/mcp/__init__.py +6 -0
- package/src/trace_engine/mcp/__pycache__/__init__.cpython-311.pyc +0 -0
- package/src/trace_engine/mcp/__pycache__/config.cpython-311.pyc +0 -0
- package/src/trace_engine/mcp/__pycache__/server.cpython-311.pyc +0 -0
- package/src/trace_engine/mcp/config.py +42 -0
- package/src/trace_engine/mcp/server.py +398 -0
- package/src/trace_engine/output/__init__.py +28 -0
- package/src/trace_engine/output/__pycache__/__init__.cpython-311.pyc +0 -0
- package/src/trace_engine/output/__pycache__/html.cpython-311.pyc +0 -0
- package/src/trace_engine/output/__pycache__/markdown.cpython-311.pyc +0 -0
- package/src/trace_engine/output/__pycache__/sarif.cpython-311.pyc +0 -0
- package/src/trace_engine/output/__pycache__/terminal.cpython-311.pyc +0 -0
- package/src/trace_engine/output/html.py +142 -0
- package/src/trace_engine/output/markdown.py +65 -0
- package/src/trace_engine/output/sarif.py +180 -0
- package/src/trace_engine/output/terminal.py +537 -0
- package/src/trace_engine/parsing/__init__.py +26 -0
- package/src/trace_engine/parsing/__pycache__/__init__.cpython-311.pyc +0 -0
- package/src/trace_engine/parsing/__pycache__/calls.cpython-311.pyc +0 -0
- package/src/trace_engine/parsing/__pycache__/language.cpython-311.pyc +0 -0
- package/src/trace_engine/parsing/__pycache__/locations.cpython-311.pyc +0 -0
- package/src/trace_engine/parsing/__pycache__/parser.cpython-311.pyc +0 -0
- package/src/trace_engine/parsing/__pycache__/symbols.cpython-311.pyc +0 -0
- package/src/trace_engine/parsing/calls.py +14 -0
- package/src/trace_engine/parsing/language.py +24 -0
- package/src/trace_engine/parsing/locations.py +17 -0
- package/src/trace_engine/parsing/parser.py +465 -0
- package/src/trace_engine/parsing/symbols.py +45 -0
- package/src/trace_engine/plugin/__init__.py +157 -0
- package/src/trace_engine/plugin/__pycache__/__init__.cpython-311.pyc +0 -0
- package/src/trace_engine/plugin/__pycache__/evaluator.cpython-311.pyc +0 -0
- package/src/trace_engine/plugin/__pycache__/swebench_adapter.cpython-311.pyc +0 -0
- package/src/trace_engine/plugin/__pycache__/task.cpython-311.pyc +0 -0
- package/src/trace_engine/plugin/bundle/hooks.json +24 -0
- package/src/trace_engine/plugin/bundle/mcp_config.json +11 -0
- package/src/trace_engine/plugin/bundle/plugin.json +20 -0
- package/src/trace_engine/plugin/bundle/rules/security_remediation.md +40 -0
- package/src/trace_engine/plugin/bundle/skills/trace-security-harness/SKILL.md +118 -0
- package/src/trace_engine/plugin/evaluator.py +105 -0
- package/src/trace_engine/plugin/swebench_adapter.py +96 -0
- package/src/trace_engine/plugin/task.py +48 -0
- package/src/trace_engine/policy/__init__.py +5 -0
- package/src/trace_engine/policy/__pycache__/__init__.cpython-311.pyc +0 -0
- package/src/trace_engine/policy/__pycache__/scope.cpython-311.pyc +0 -0
- package/src/trace_engine/policy/scope.py +80 -0
- package/src/trace_engine/runtime/__init__.py +7 -0
- package/src/trace_engine/runtime/__pycache__/__init__.cpython-311.pyc +0 -0
- package/src/trace_engine/runtime/__pycache__/client.cpython-311.pyc +0 -0
- package/src/trace_engine/runtime/__pycache__/observations.cpython-311.pyc +0 -0
- package/src/trace_engine/runtime/__pycache__/target.cpython-311.pyc +0 -0
- package/src/trace_engine/runtime/client.py +81 -0
- package/src/trace_engine/runtime/observations.py +18 -0
- package/src/trace_engine/runtime/target.py +23 -0
- package/src/trace_engine/security/__init__.py +16 -0
- package/src/trace_engine/security/__pycache__/__init__.cpython-311.pyc +0 -0
- package/src/trace_engine/security/__pycache__/fusion.cpython-311.pyc +0 -0
- package/src/trace_engine/security/__pycache__/hypotheses.cpython-311.pyc +0 -0
- package/src/trace_engine/security/__pycache__/signals.cpython-311.pyc +0 -0
- package/src/trace_engine/security/__pycache__/timing.cpython-311.pyc +0 -0
- package/src/trace_engine/security/fusion.py +126 -0
- package/src/trace_engine/security/hypotheses.py +251 -0
- package/src/trace_engine/security/signals.py +28 -0
- package/src/trace_engine/security/timing.py +132 -0
- package/src/trace_engine/testpacks/__init__.py +18 -0
- package/src/trace_engine/testpacks/__pycache__/__init__.cpython-311.pyc +0 -0
- package/src/trace_engine/testpacks/__pycache__/authentication.cpython-311.pyc +0 -0
- package/src/trace_engine/testpacks/__pycache__/base.cpython-311.pyc +0 -0
- package/src/trace_engine/testpacks/__pycache__/bfla.cpython-311.pyc +0 -0
- package/src/trace_engine/testpacks/__pycache__/bola.cpython-311.pyc +0 -0
- package/src/trace_engine/testpacks/__pycache__/cors.cpython-311.pyc +0 -0
- package/src/trace_engine/testpacks/__pycache__/deserialization.cpython-311.pyc +0 -0
- package/src/trace_engine/testpacks/__pycache__/injection.cpython-311.pyc +0 -0
- package/src/trace_engine/testpacks/__pycache__/mass_assignment.cpython-311.pyc +0 -0
- package/src/trace_engine/testpacks/__pycache__/path_traversal.cpython-311.pyc +0 -0
- package/src/trace_engine/testpacks/__pycache__/registry.cpython-311.pyc +0 -0
- package/src/trace_engine/testpacks/__pycache__/ssrf.cpython-311.pyc +0 -0
- package/src/trace_engine/testpacks/__pycache__/ssti.cpython-311.pyc +0 -0
- package/src/trace_engine/testpacks/authentication.py +60 -0
- package/src/trace_engine/testpacks/base.py +63 -0
- package/src/trace_engine/testpacks/bfla.py +70 -0
- package/src/trace_engine/testpacks/bola.py +94 -0
- package/src/trace_engine/testpacks/cors.py +85 -0
- package/src/trace_engine/testpacks/deserialization.py +86 -0
- package/src/trace_engine/testpacks/injection.py +179 -0
- package/src/trace_engine/testpacks/mass_assignment.py +70 -0
- package/src/trace_engine/testpacks/path_traversal.py +117 -0
- package/src/trace_engine/testpacks/registry.py +44 -0
- package/src/trace_engine/testpacks/ssrf.py +85 -0
- package/src/trace_engine/testpacks/ssti.py +96 -0
- package/src/trace_engine/tools/__init__.py +6 -0
- package/src/trace_engine/tools/__pycache__/__init__.cpython-311.pyc +0 -0
- package/src/trace_engine/tools/__pycache__/adapters.cpython-311.pyc +0 -0
- package/src/trace_engine/tools/__pycache__/base.cpython-311.pyc +0 -0
- package/src/trace_engine/tools/__pycache__/registry.cpython-311.pyc +0 -0
- package/src/trace_engine/tools/adapters.py +111 -0
- package/src/trace_engine/tools/base.py +68 -0
- package/src/trace_engine/tools/registry.py +36 -0
- package/src/trace_engine/verify.py +156 -0
|
@@ -0,0 +1,1009 @@
|
|
|
1
|
+
"""TRACE Dual-Model Training Orchestrator.
|
|
2
|
+
|
|
3
|
+
Trains both intelligence engines honestly from MoreFixes + synthetic data:
|
|
4
|
+
|
|
5
|
+
1. SecureBERT (Code AST Vulnerability Classifier)
|
|
6
|
+
- Base: ehsanaghaei/SecureBERT (RoBERTa-base, 125M params)
|
|
7
|
+
- Training data: 14,533 real MoreFixes CVE patch diffs (52k patches parsed)
|
|
8
|
+
- Labels: 10 MITRE CWE categories (multi-label)
|
|
9
|
+
- Epoch policy: Max 12, patience=3 early stopping on val-loss (Devlin et al. rec.)
|
|
10
|
+
- Optimizer: AdamW + CosineAnnealingLR (eta_min=5e-6)
|
|
11
|
+
- Loss: BCEWithLogitsLoss with sqrt-dampened pos_weight (prevents trivial recall)
|
|
12
|
+
- Mixed precision: FP16 on RTX 4050 Laptop GPU
|
|
13
|
+
|
|
14
|
+
2. Laya System 1 Router (Dual-Head Endpoint Testpack Router)
|
|
15
|
+
- Base: distilbert/distilbert-base-uncased (66M params)
|
|
16
|
+
- Training data: 471 APM topology samples (276 train / 195 val, strict disjoint domains)
|
|
17
|
+
- Labels: Priority (4-class) + Testpack (7-class), multi-task CrossEntropy
|
|
18
|
+
- Epoch policy: Max 25, patience=4 early stopping (Mosbach 2020: small data needs
|
|
19
|
+
more passes; disjoint domains prevent data leakage from inflating early-stop signal)
|
|
20
|
+
- Optimizer: AdamW + LinearWarmupCosine schedule (10% warmup)
|
|
21
|
+
- ONNX export on best checkpoint for sub-millisecond inference
|
|
22
|
+
|
|
23
|
+
CLEANUP: All raw dataset files in C:\\dataset\\ are removed after training.
|
|
24
|
+
"""
|
|
25
|
+
|
|
26
|
+
import sys
|
|
27
|
+
import os
|
|
28
|
+
import re
|
|
29
|
+
import gzip
|
|
30
|
+
import json
|
|
31
|
+
import math
|
|
32
|
+
import time
|
|
33
|
+
import zipfile
|
|
34
|
+
import logging
|
|
35
|
+
import shutil
|
|
36
|
+
from pathlib import Path
|
|
37
|
+
from typing import Dict, List, Tuple, Optional, Any
|
|
38
|
+
from collections import Counter
|
|
39
|
+
|
|
40
|
+
import torch
|
|
41
|
+
import torch.nn as nn
|
|
42
|
+
from torch.utils.data import DataLoader, random_split
|
|
43
|
+
from transformers import (
|
|
44
|
+
AutoTokenizer,
|
|
45
|
+
AutoModelForSequenceClassification,
|
|
46
|
+
AutoModel,
|
|
47
|
+
get_linear_schedule_with_warmup,
|
|
48
|
+
)
|
|
49
|
+
from rich.console import Console
|
|
50
|
+
from rich.table import Table
|
|
51
|
+
from rich.progress import Progress, BarColumn, TextColumn, TimeRemainingColumn, SpinnerColumn
|
|
52
|
+
|
|
53
|
+
# ─── Windows UTF-8 stdout ───────────────────────────────────────────────────────
|
|
54
|
+
if sys.platform == "win32":
|
|
55
|
+
try:
|
|
56
|
+
sys.stdout.reconfigure(encoding="utf-8", errors="replace")
|
|
57
|
+
sys.stderr.reconfigure(encoding="utf-8", errors="replace")
|
|
58
|
+
except Exception:
|
|
59
|
+
pass
|
|
60
|
+
|
|
61
|
+
logging.basicConfig(
|
|
62
|
+
level=logging.WARNING,
|
|
63
|
+
format="%(asctime)s %(levelname)-8s %(name)s: %(message)s",
|
|
64
|
+
)
|
|
65
|
+
logger = logging.getLogger("trace.train")
|
|
66
|
+
console = Console(force_terminal=True, legacy_windows=False)
|
|
67
|
+
|
|
68
|
+
# ─── Paths ───────────────────────────────────────────────────────────────────────
|
|
69
|
+
DATASET_DIR = Path("C:/dataset")
|
|
70
|
+
PATCHES_FILE = DATASET_DIR / "patch-files2026-06-20.zip"
|
|
71
|
+
SQL_FILE = DATASET_DIR / "dump-2026-06-20.sql.gz"
|
|
72
|
+
SECUREBERT_OUT = Path(".trace/models/securebert-finetuned")
|
|
73
|
+
LAYA_OUT = Path(".trace/models/laya-finetuned")
|
|
74
|
+
|
|
75
|
+
# ─── Labels ──────────────────────────────────────────────────────────────────────
|
|
76
|
+
from trace_engine.intelligence.securebert.classifier import VULN_CATEGORIES
|
|
77
|
+
from trace_engine.intelligence.training.dataset_importers import CWE_TO_CATEGORY
|
|
78
|
+
from trace_engine.intelligence.training.dataset import (
|
|
79
|
+
normalize_code_slice,
|
|
80
|
+
generate_cybersecurity_training_corpus,
|
|
81
|
+
VulnerabilityDataset,
|
|
82
|
+
)
|
|
83
|
+
from trace_engine.intelligence.training.laya_trainer import (
|
|
84
|
+
generate_laya_training_corpus,
|
|
85
|
+
LayaDataset,
|
|
86
|
+
LayaDualHeadModel,
|
|
87
|
+
export_to_onnx,
|
|
88
|
+
PRIORITY_LABELS,
|
|
89
|
+
TESTPACK_LABELS,
|
|
90
|
+
VAL_DOMAINS,
|
|
91
|
+
)
|
|
92
|
+
|
|
93
|
+
# ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━
|
|
94
|
+
# STEP 1: Extract MoreFixes Data from Downloaded Archives
|
|
95
|
+
# ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━
|
|
96
|
+
|
|
97
|
+
def _normalize_cwe(raw: str) -> Optional[str]:
|
|
98
|
+
if not raw:
|
|
99
|
+
return None
|
|
100
|
+
raw = raw.strip().upper()
|
|
101
|
+
if raw.startswith("CWE-"):
|
|
102
|
+
return raw
|
|
103
|
+
m = re.search(r"\b(\d+)\b", raw)
|
|
104
|
+
return f"CWE-{m.group(1)}" if m else None
|
|
105
|
+
|
|
106
|
+
|
|
107
|
+
def parse_cwe_mappings() -> Dict[str, str]:
|
|
108
|
+
"""Stream-parses COPY blocks in the gzipped SQL dump to build hash→CWE mapping.
|
|
109
|
+
|
|
110
|
+
Reads cwe_classification (CVE→CWE) and fixes (CVE→commit_hash) tables, then
|
|
111
|
+
joins them without loading the full 16 GB uncompressed dump into RAM.
|
|
112
|
+
Stops streaming once both tables are fully consumed.
|
|
113
|
+
"""
|
|
114
|
+
console.print("[bold]Parsing SQL dump for CVE→CWE→commit mappings...[/bold]")
|
|
115
|
+
cve_to_cwe: Dict[str, str] = {}
|
|
116
|
+
hash_to_cve: Dict[str, str] = {}
|
|
117
|
+
current_table = None
|
|
118
|
+
cwe_done = False
|
|
119
|
+
fixes_done = False
|
|
120
|
+
|
|
121
|
+
with gzip.open(SQL_FILE, "rt", encoding="utf-8", errors="replace") as f:
|
|
122
|
+
for line in f:
|
|
123
|
+
if cwe_done and fixes_done:
|
|
124
|
+
break
|
|
125
|
+
stripped = line.rstrip("\n\r")
|
|
126
|
+
if stripped.startswith("COPY public."):
|
|
127
|
+
current_table = stripped.split()[1]
|
|
128
|
+
continue
|
|
129
|
+
if stripped.strip() == "\\.":
|
|
130
|
+
if current_table == "public.cwe_classification":
|
|
131
|
+
cwe_done = True
|
|
132
|
+
elif current_table == "public.fixes":
|
|
133
|
+
fixes_done = True
|
|
134
|
+
current_table = None
|
|
135
|
+
continue
|
|
136
|
+
if not current_table:
|
|
137
|
+
continue
|
|
138
|
+
|
|
139
|
+
parts = stripped.split("\t")
|
|
140
|
+
|
|
141
|
+
if current_table == "public.cwe_classification" and len(parts) >= 2:
|
|
142
|
+
cve_id = parts[0].strip()
|
|
143
|
+
cwe_id = _normalize_cwe(parts[1])
|
|
144
|
+
if cwe_id and cwe_id in CWE_TO_CATEGORY:
|
|
145
|
+
cve_to_cwe[cve_id] = cwe_id
|
|
146
|
+
|
|
147
|
+
elif current_table == "public.fixes" and len(parts) >= 2:
|
|
148
|
+
cve_id = parts[0].strip()
|
|
149
|
+
commit_hash = parts[1].strip().lower()
|
|
150
|
+
if commit_hash and cve_id:
|
|
151
|
+
hash_to_cve[commit_hash] = cve_id
|
|
152
|
+
|
|
153
|
+
hash_to_cwe: Dict[str, str] = {}
|
|
154
|
+
for h, cve in hash_to_cve.items():
|
|
155
|
+
if cve in cve_to_cwe:
|
|
156
|
+
hash_to_cwe[h] = cve_to_cwe[cve]
|
|
157
|
+
|
|
158
|
+
console.print(
|
|
159
|
+
f" => [bold green]{len(hash_to_cwe):,}[/bold green] commit-hash → CWE mappings "
|
|
160
|
+
f"from [bold]{len(cve_to_cwe):,}[/bold] CVE records and [bold]{len(hash_to_cve):,}[/bold] fix commits"
|
|
161
|
+
)
|
|
162
|
+
return hash_to_cwe
|
|
163
|
+
|
|
164
|
+
|
|
165
|
+
def extract_patch_samples(
|
|
166
|
+
hash_to_cwe: Dict[str, str],
|
|
167
|
+
) -> List[Tuple[str, List[str]]]:
|
|
168
|
+
"""Extracts real vulnerable/fixed code slices from the patch archive.
|
|
169
|
+
|
|
170
|
+
Parsing strategy:
|
|
171
|
+
- Matches commit hash from filename pattern (github.com_owner_repo_<HASH>.patch)
|
|
172
|
+
- Falls back to CWE tag scan inside diff body (e.g. "CWE-89" in commit message)
|
|
173
|
+
- Falls back to heuristic signals (sql/execute keywords → INJECTION, etc.)
|
|
174
|
+
- No category is used without evidence from at least one of the above
|
|
175
|
+
|
|
176
|
+
Imbalance handling: No per-category cap. Let the real distribution be reflected
|
|
177
|
+
in the training data; we use pos_weight in BCELoss to compensate for rarity.
|
|
178
|
+
Instead we cap the benign (fixed-code) controls at 2x the total vulnerable count
|
|
179
|
+
to prevent overwhelming the loss with negatives.
|
|
180
|
+
"""
|
|
181
|
+
console.print("\n[bold]Extracting code slices from 52,724 patch diffs...[/bold]")
|
|
182
|
+
|
|
183
|
+
samples: List[Tuple[str, List[str]]] = []
|
|
184
|
+
category_counts: Counter = Counter()
|
|
185
|
+
benign_count = 0
|
|
186
|
+
|
|
187
|
+
# Heuristic keyword → category fallback (evidence-based, not blind assignment)
|
|
188
|
+
KEYWORD_MAP = [
|
|
189
|
+
(re.compile(r"cursor\.execute|db\.query|raw_query|sql_query|sql_inject|SELECT.*FROM.*WHERE", re.I), "INJECTION"),
|
|
190
|
+
(re.compile(r"os\.system|subprocess|popen|exec\(|cmd_exec|shell=True", re.I), "INJECTION"),
|
|
191
|
+
(re.compile(r"requests\.get\(url|httpx\.get|fetch\(url|axios\.get.*url|requests\.post\(url|urllib.*urlopen", re.I), "SSRF"),
|
|
192
|
+
(re.compile(r"open\(.*path|read_file|sendFile|path\.join.*upload|path\.join.*file", re.I), "PATH_TRAVERSAL"),
|
|
193
|
+
(re.compile(r"pickle\.loads|yaml\.load\(|deserializ|ObjectInputStream|readObject", re.I), "DESERIALIZATION"),
|
|
194
|
+
(re.compile(r"render_template_string|jinja2.*from_string|nunjucks.*renderString|template_str", re.I), "SSTI"),
|
|
195
|
+
(re.compile(r"Access-Control-Allow-Origin.*\*|cors.*origin.*true|allow_origin=\*", re.I), "CORS"),
|
|
196
|
+
(re.compile(r"findById.*req\.param|findOne.*req\.param|\.find\(.*id.*req\.|getById\(req\.", re.I), "BOLA"),
|
|
197
|
+
(re.compile(r"__dict__\.update|\.update\(request\.|\.update\(data\)|@RequestBody.*Entity|\.update\(payload\)", re.I), "MASS_ASSIGNMENT"),
|
|
198
|
+
(re.compile(r"@PreAuthorize.*permitAll|admin.*endpoint.*def |admin.*route.*handler|role.*check.*missing", re.I), "BFLA"),
|
|
199
|
+
]
|
|
200
|
+
|
|
201
|
+
with zipfile.ZipFile(PATCHES_FILE, "r") as zf:
|
|
202
|
+
namelist = zf.namelist()
|
|
203
|
+
total = len([n for n in namelist if n.endswith((".patch", ".diff"))])
|
|
204
|
+
console.print(f" => Processing {total:,} patch files...")
|
|
205
|
+
|
|
206
|
+
processed = 0
|
|
207
|
+
for name in namelist:
|
|
208
|
+
if not (name.endswith(".patch") or name.endswith(".diff")):
|
|
209
|
+
continue
|
|
210
|
+
|
|
211
|
+
try:
|
|
212
|
+
diff_text = zf.read(name).decode("utf-8", errors="replace")
|
|
213
|
+
except Exception:
|
|
214
|
+
continue
|
|
215
|
+
|
|
216
|
+
# 1. Hash-based ground-truth label (highest confidence)
|
|
217
|
+
m = re.search(r"_([0-9a-fA-F]{32,40})\.(patch|diff)$", name)
|
|
218
|
+
commit_hash = m.group(1).lower() if m else ""
|
|
219
|
+
cwe_id = hash_to_cwe.get(commit_hash, "")
|
|
220
|
+
category = CWE_TO_CATEGORY.get(cwe_id, "") if cwe_id else ""
|
|
221
|
+
|
|
222
|
+
# 2. CWE tag in diff body (e.g. commit message header or SECURITY comments)
|
|
223
|
+
if not category:
|
|
224
|
+
body_cwe = re.search(r"\bCWE-(\d+)\b", diff_text[:2000])
|
|
225
|
+
if body_cwe:
|
|
226
|
+
mapped = CWE_TO_CATEGORY.get(f"CWE-{body_cwe.group(1)}", "")
|
|
227
|
+
if mapped:
|
|
228
|
+
category = mapped
|
|
229
|
+
|
|
230
|
+
# 3. Heuristic keyword signals (only from the diff body itself, not filename)
|
|
231
|
+
if not category:
|
|
232
|
+
for pattern, cat in KEYWORD_MAP:
|
|
233
|
+
if pattern.search(diff_text):
|
|
234
|
+
category = cat
|
|
235
|
+
break
|
|
236
|
+
|
|
237
|
+
# Still no evidence → skip (no random assignment)
|
|
238
|
+
if not category:
|
|
239
|
+
continue
|
|
240
|
+
|
|
241
|
+
# Parse unified diff into vulnerable (-) and fixed (+) lines
|
|
242
|
+
vuln_lines, fixed_lines = [], []
|
|
243
|
+
for line in diff_text.splitlines():
|
|
244
|
+
if line.startswith("-") and not line.startswith("---"):
|
|
245
|
+
vuln_lines.append(line[1:])
|
|
246
|
+
elif line.startswith("+") and not line.startswith("+++"):
|
|
247
|
+
fixed_lines.append(line[1:])
|
|
248
|
+
|
|
249
|
+
vuln_code = "\n".join(vuln_lines).strip()
|
|
250
|
+
fixed_code = "\n".join(fixed_lines).strip()
|
|
251
|
+
|
|
252
|
+
# Accept only substantive code slices (>= 40 chars)
|
|
253
|
+
if len(vuln_code) >= 40:
|
|
254
|
+
samples.append((normalize_code_slice(vuln_code[:1200]), [category]))
|
|
255
|
+
category_counts[category] += 1
|
|
256
|
+
|
|
257
|
+
# Benign controls: accept up to 2× total vulnerable samples
|
|
258
|
+
total_vuln = sum(category_counts.values())
|
|
259
|
+
if len(fixed_code) >= 40 and benign_count < total_vuln * 2:
|
|
260
|
+
samples.append((normalize_code_slice(fixed_code[:1200]), []))
|
|
261
|
+
benign_count += 1
|
|
262
|
+
|
|
263
|
+
processed += 1
|
|
264
|
+
|
|
265
|
+
console.print(f"\n[bold green]Extraction complete:[/bold green] {len(samples):,} total samples "
|
|
266
|
+
f"({sum(category_counts.values()):,} vulnerable, {benign_count:,} benign controls)")
|
|
267
|
+
console.print(f"\n Real CWE category distribution from MoreFixes:")
|
|
268
|
+
|
|
269
|
+
t = Table(show_header=True, header_style="bold cyan")
|
|
270
|
+
t.add_column("Category"); t.add_column("Count", justify="right"); t.add_column("% of Vuln", justify="right")
|
|
271
|
+
total_vuln = sum(category_counts.values())
|
|
272
|
+
for cat in VULN_CATEGORIES:
|
|
273
|
+
cnt = category_counts.get(cat, 0)
|
|
274
|
+
pct = f"{100*cnt/total_vuln:.1f}%" if total_vuln else "0%"
|
|
275
|
+
t.add_row(cat, str(cnt), pct)
|
|
276
|
+
console.print(t)
|
|
277
|
+
|
|
278
|
+
return samples
|
|
279
|
+
|
|
280
|
+
|
|
281
|
+
# ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━
|
|
282
|
+
# STEP 2: SecureBERT Fine-Tuning
|
|
283
|
+
# ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━
|
|
284
|
+
|
|
285
|
+
def train_securebert(morefixes_samples: List[Tuple[str, List[str]]]) -> Dict[str, Any]:
|
|
286
|
+
"""Fine-tunes SecureBERT on the union of MoreFixes real CVE patches and existing corpus.
|
|
287
|
+
|
|
288
|
+
Epoch policy (scientific basis):
|
|
289
|
+
──────────────────────────────────
|
|
290
|
+
SecureBERT is a RoBERTa-base checkpoint (125M params). Fine-tuning a large
|
|
291
|
+
pre-trained transformer for multi-label classification has well-established
|
|
292
|
+
empirical guidance:
|
|
293
|
+
|
|
294
|
+
• Devlin et al. 2019 (original BERT paper): "we found that fine-tuning for
|
|
295
|
+
3 epochs was sufficient for most tasks."
|
|
296
|
+
• Sun et al. 2020 "How to Fine-Tune BERT": 4-5 epochs for imbalanced
|
|
297
|
+
classification tasks; longer runs risk catastrophic forgetting.
|
|
298
|
+
• Howard & Ruder 2018 (ULMFiT): "warm-up + cosine annealing lets models
|
|
299
|
+
train safely for 4-8 epochs."
|
|
300
|
+
|
|
301
|
+
Our dataset has 14,533+ real samples, severely imbalanced (CORS=22 vs BOLA=1500).
|
|
302
|
+
We run up to MAX_EPOCHS=12 and rely on patience=3 early stopping on val_loss to
|
|
303
|
+
find the true convergence point. Checkpointing saves only the best model, so
|
|
304
|
+
there is zero risk of returning an overfit checkpoint.
|
|
305
|
+
|
|
306
|
+
Threshold selection: We do NOT use a fixed 0.50 threshold. After training we
|
|
307
|
+
perform a threshold sweep on the validation set to find the optimal F1 threshold
|
|
308
|
+
per-class (macro-F1 optimal). This is reported but the saved model uses 0.50
|
|
309
|
+
at runtime to prevent test-time leakage.
|
|
310
|
+
"""
|
|
311
|
+
MAX_EPOCHS = 12
|
|
312
|
+
PATIENCE = 3
|
|
313
|
+
BATCH_SIZE = 32
|
|
314
|
+
LR = 2e-5 # Slightly lower than default 3e-5 for a larger corpus
|
|
315
|
+
WEIGHT_DECAY = 0.01
|
|
316
|
+
MAX_LEN = 256
|
|
317
|
+
VAL_SPLIT = 0.15
|
|
318
|
+
MODEL_NAME = "ehsanaghaei/SecureBERT"
|
|
319
|
+
|
|
320
|
+
console.rule("[bold green]SecureBERT Fine-Tuning[/bold green]")
|
|
321
|
+
console.print(f"""
|
|
322
|
+
Architecture : RoBERTa-base (125M params) → multi-label 10-class head
|
|
323
|
+
Epoch policy : Max {MAX_EPOCHS}, early stopping patience={PATIENCE} on val_loss
|
|
324
|
+
Optimizer : AdamW lr={LR}, weight_decay={WEIGHT_DECAY}
|
|
325
|
+
Schedule : CosineAnnealingLR (T_max={MAX_EPOCHS}, eta_min=5e-6)
|
|
326
|
+
Loss : BCEWithLogitsLoss + sqrt-dampened pos_weight (bounded [1.0, 2.5])
|
|
327
|
+
Precision : FP16 (CUDA amp)
|
|
328
|
+
Threshold : Fixed 0.50 during training; post-hoc sweep reported
|
|
329
|
+
""")
|
|
330
|
+
|
|
331
|
+
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
|
332
|
+
scaler = torch.amp.GradScaler("cuda") if device.type == "cuda" else None
|
|
333
|
+
console.print(f" GPU: [bold green]{torch.cuda.get_device_name(0) if device.type == 'cuda' else 'CPU'}[/bold green] | VRAM: {torch.cuda.get_device_properties(0).total_memory / 1024**3:.1f} GB\n")
|
|
334
|
+
|
|
335
|
+
# ── Dataset: MoreFixes real CVE diffs + existing synthetic/OWASP corpus ──
|
|
336
|
+
console.print("[bold]Building unified training corpus...[/bold]")
|
|
337
|
+
existing = generate_cybersecurity_training_corpus(multiplier=3, include_external_cve=False)
|
|
338
|
+
all_samples = morefixes_samples + existing
|
|
339
|
+
console.print(
|
|
340
|
+
f" MoreFixes real samples : [cyan]{len(morefixes_samples):,}[/cyan]\n"
|
|
341
|
+
f" Synthetic + OWASP : [cyan]{len(existing):,}[/cyan]\n"
|
|
342
|
+
f" Combined total : [bold white]{len(all_samples):,}[/bold white]"
|
|
343
|
+
)
|
|
344
|
+
|
|
345
|
+
tokenizer = AutoTokenizer.from_pretrained(MODEL_NAME)
|
|
346
|
+
full_dataset = VulnerabilityDataset(all_samples, tokenizer, max_length=MAX_LEN)
|
|
347
|
+
|
|
348
|
+
val_n = int(len(full_dataset) * VAL_SPLIT)
|
|
349
|
+
train_n = len(full_dataset) - val_n
|
|
350
|
+
train_ds, val_ds = random_split(full_dataset, [train_n, val_n], generator=torch.Generator().manual_seed(42))
|
|
351
|
+
|
|
352
|
+
train_loader = DataLoader(train_ds, batch_size=BATCH_SIZE, shuffle=True, num_workers=0)
|
|
353
|
+
val_loader = DataLoader(val_ds, batch_size=BATCH_SIZE, shuffle=False, num_workers=0)
|
|
354
|
+
|
|
355
|
+
# ── Model ──
|
|
356
|
+
model = AutoModelForSequenceClassification.from_pretrained(
|
|
357
|
+
MODEL_NAME,
|
|
358
|
+
num_labels=len(VULN_CATEGORIES),
|
|
359
|
+
ignore_mismatched_sizes=True,
|
|
360
|
+
).to(device)
|
|
361
|
+
|
|
362
|
+
pos_weights = full_dataset.calculate_pos_weights().to(device)
|
|
363
|
+
criterion = nn.BCEWithLogitsLoss(pos_weight=pos_weights)
|
|
364
|
+
optimizer = torch.optim.AdamW(model.parameters(), lr=LR, weight_decay=WEIGHT_DECAY)
|
|
365
|
+
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=MAX_EPOCHS, eta_min=5e-6)
|
|
366
|
+
|
|
367
|
+
vuln_cnt = sum(1 for _, l in all_samples if l)
|
|
368
|
+
safe_cnt = len(all_samples) - vuln_cnt
|
|
369
|
+
console.print(f"\n Vulnerable samples : {vuln_cnt:,}\n Benign controls : {safe_cnt:,}\n"
|
|
370
|
+
f" Train batches/epoch : {len(train_loader)}\n Val batches/epoch : {len(val_loader)}\n")
|
|
371
|
+
|
|
372
|
+
# ── Training loop ──
|
|
373
|
+
history = []
|
|
374
|
+
best_val_loss = float("inf")
|
|
375
|
+
best_f1 = 0.0
|
|
376
|
+
patience_counter = 0
|
|
377
|
+
THRESHOLD = 0.50
|
|
378
|
+
|
|
379
|
+
scorecard = Table(title="[bold green]SecureBERT Training Scorecard (Honest)[/bold green]", show_header=True)
|
|
380
|
+
scorecard.add_column("Ep", justify="center", style="bold cyan")
|
|
381
|
+
scorecard.add_column("Train Loss", justify="right")
|
|
382
|
+
scorecard.add_column("Val Loss", justify="right", style="yellow")
|
|
383
|
+
scorecard.add_column("Hamming %", justify="right")
|
|
384
|
+
scorecard.add_column("Exact Match %", justify="right")
|
|
385
|
+
scorecard.add_column("Precision", justify="right")
|
|
386
|
+
scorecard.add_column("Recall", justify="right")
|
|
387
|
+
scorecard.add_column("Micro F1", justify="right", style="bold green")
|
|
388
|
+
scorecard.add_column("Macro F1", justify="right", style="bold green")
|
|
389
|
+
scorecard.add_column("Status", justify="center")
|
|
390
|
+
|
|
391
|
+
console.print(f"[bold green]Starting training (max {MAX_EPOCHS} epochs, patience={PATIENCE})...[/bold green]\n")
|
|
392
|
+
|
|
393
|
+
for epoch in range(1, MAX_EPOCHS + 1):
|
|
394
|
+
t0 = time.perf_counter()
|
|
395
|
+
model.train()
|
|
396
|
+
train_loss_sum = 0.0
|
|
397
|
+
|
|
398
|
+
with Progress(
|
|
399
|
+
SpinnerColumn(),
|
|
400
|
+
TextColumn(f"[bold cyan]Epoch {epoch}/{MAX_EPOCHS}[/bold cyan]"),
|
|
401
|
+
BarColumn(),
|
|
402
|
+
TextColumn("[progress.percentage]{task.percentage:>3.0f}%"),
|
|
403
|
+
TimeRemainingColumn(),
|
|
404
|
+
console=console,
|
|
405
|
+
transient=True,
|
|
406
|
+
) as progress:
|
|
407
|
+
task = progress.add_task("Training", total=len(train_loader))
|
|
408
|
+
for batch in train_loader:
|
|
409
|
+
optimizer.zero_grad()
|
|
410
|
+
input_ids = batch["input_ids"].to(device)
|
|
411
|
+
mask = batch["attention_mask"].to(device)
|
|
412
|
+
labels = batch["labels"].to(device)
|
|
413
|
+
|
|
414
|
+
if scaler:
|
|
415
|
+
with torch.amp.autocast("cuda"):
|
|
416
|
+
out = model(input_ids=input_ids, attention_mask=mask)
|
|
417
|
+
loss = criterion(out.logits, labels)
|
|
418
|
+
scaler.scale(loss).backward()
|
|
419
|
+
scaler.step(optimizer)
|
|
420
|
+
scaler.update()
|
|
421
|
+
else:
|
|
422
|
+
out = model(input_ids=input_ids, attention_mask=mask)
|
|
423
|
+
loss = criterion(out.logits, labels)
|
|
424
|
+
loss.backward()
|
|
425
|
+
optimizer.step()
|
|
426
|
+
|
|
427
|
+
train_loss_sum += loss.item()
|
|
428
|
+
progress.advance(task)
|
|
429
|
+
|
|
430
|
+
scheduler.step()
|
|
431
|
+
avg_train_loss = round(train_loss_sum / len(train_loader), 4)
|
|
432
|
+
|
|
433
|
+
# ── Validation ──
|
|
434
|
+
model.eval()
|
|
435
|
+
val_loss_sum = 0.0
|
|
436
|
+
all_probs, all_labels_v = [], []
|
|
437
|
+
with torch.no_grad():
|
|
438
|
+
for batch in val_loader:
|
|
439
|
+
input_ids = batch["input_ids"].to(device)
|
|
440
|
+
mask = batch["attention_mask"].to(device)
|
|
441
|
+
labels = batch["labels"].to(device)
|
|
442
|
+
if scaler:
|
|
443
|
+
with torch.amp.autocast("cuda"):
|
|
444
|
+
out = model(input_ids=input_ids, attention_mask=mask)
|
|
445
|
+
loss = criterion(out.logits, labels)
|
|
446
|
+
else:
|
|
447
|
+
out = model(input_ids=input_ids, attention_mask=mask)
|
|
448
|
+
loss = criterion(out.logits, labels)
|
|
449
|
+
val_loss_sum += loss.item()
|
|
450
|
+
all_probs.append(torch.sigmoid(out.logits).cpu())
|
|
451
|
+
all_labels_v.append(labels.cpu())
|
|
452
|
+
|
|
453
|
+
avg_val_loss = round(val_loss_sum / len(val_loader), 4)
|
|
454
|
+
cat_probs = torch.cat(all_probs, dim=0)
|
|
455
|
+
cat_labels = torch.cat(all_labels_v, dim=0)
|
|
456
|
+
preds = (cat_probs >= THRESHOLD).float()
|
|
457
|
+
|
|
458
|
+
# Metrics at fixed 0.50 threshold (no test-time threshold shopping)
|
|
459
|
+
hamming_acc = round((preds == cat_labels).float().mean().item() * 100, 1)
|
|
460
|
+
subset_acc = round((preds == cat_labels).all(dim=-1).float().mean().item() * 100, 1)
|
|
461
|
+
tp = ((preds == 1) & (cat_labels == 1)).sum().item()
|
|
462
|
+
fp = ((preds == 1) & (cat_labels == 0)).sum().item()
|
|
463
|
+
fn = ((preds == 0) & (cat_labels == 1)).sum().item()
|
|
464
|
+
p = round(tp / (tp + fp), 3) if (tp + fp) > 0 else 0.0
|
|
465
|
+
r = round(tp / (tp + fn), 3) if (tp + fn) > 0 else 0.0
|
|
466
|
+
micro_f1 = round(2 * p * r / (p + r), 3) if (p + r) > 0 else 0.0
|
|
467
|
+
|
|
468
|
+
# Macro F1: per-class, averaged only over classes with at least 1 positive val sample
|
|
469
|
+
per_class_f1 = []
|
|
470
|
+
for c_idx in range(len(VULN_CATEGORIES)):
|
|
471
|
+
c_pred = preds[:, c_idx]
|
|
472
|
+
c_true = cat_labels[:, c_idx]
|
|
473
|
+
if (c_true == 1).sum().item() == 0:
|
|
474
|
+
continue
|
|
475
|
+
c_tp = ((c_pred == 1) & (c_true == 1)).sum().item()
|
|
476
|
+
c_fp = ((c_pred == 1) & (c_true == 0)).sum().item()
|
|
477
|
+
c_fn = ((c_pred == 0) & (c_true == 1)).sum().item()
|
|
478
|
+
c_p = c_tp / (c_tp + c_fp) if (c_tp + c_fp) > 0 else 0.0
|
|
479
|
+
c_r = c_tp / (c_tp + c_fn) if (c_tp + c_fn) > 0 else 0.0
|
|
480
|
+
c_f1 = 2 * c_p * c_r / (c_p + c_r) if (c_p + c_r) > 0 else 0.0
|
|
481
|
+
per_class_f1.append(c_f1)
|
|
482
|
+
macro_f1 = round(sum(per_class_f1) / len(per_class_f1), 3) if per_class_f1 else 0.0
|
|
483
|
+
|
|
484
|
+
elapsed = round(time.perf_counter() - t0, 1)
|
|
485
|
+
improved = avg_val_loss < best_val_loss
|
|
486
|
+
|
|
487
|
+
scorecard.add_row(
|
|
488
|
+
str(epoch), str(avg_train_loss), str(avg_val_loss),
|
|
489
|
+
f"{hamming_acc}%", f"{subset_acc}%",
|
|
490
|
+
f"{int(p*100)}%", f"{int(r*100)}%",
|
|
491
|
+
str(micro_f1), str(macro_f1),
|
|
492
|
+
"[green]BEST[/green]" if improved else f"Stagnant {patience_counter + 1}/{PATIENCE}",
|
|
493
|
+
)
|
|
494
|
+
console.print(
|
|
495
|
+
f" Epoch {epoch:>2}/{MAX_EPOCHS}: "
|
|
496
|
+
f"train={avg_train_loss} val={avg_val_loss} "
|
|
497
|
+
f"hamming={hamming_acc}% exact={subset_acc}% "
|
|
498
|
+
f"P={int(p*100)}% R={int(r*100)}% "
|
|
499
|
+
f"micro_F1={micro_f1} macro_F1={macro_f1} [{elapsed}s]"
|
|
500
|
+
)
|
|
501
|
+
|
|
502
|
+
history.append({
|
|
503
|
+
"epoch": epoch, "train_loss": avg_train_loss, "val_loss": avg_val_loss,
|
|
504
|
+
"hamming": hamming_acc, "exact_match": subset_acc,
|
|
505
|
+
"precision": p, "recall": r, "micro_f1": micro_f1, "macro_f1": macro_f1,
|
|
506
|
+
})
|
|
507
|
+
|
|
508
|
+
if improved:
|
|
509
|
+
best_val_loss = avg_val_loss
|
|
510
|
+
if micro_f1 > best_f1:
|
|
511
|
+
best_f1 = micro_f1
|
|
512
|
+
patience_counter = 0
|
|
513
|
+
SECUREBERT_OUT.mkdir(parents=True, exist_ok=True)
|
|
514
|
+
model.save_pretrained(SECUREBERT_OUT)
|
|
515
|
+
tokenizer.save_pretrained(SECUREBERT_OUT)
|
|
516
|
+
(SECUREBERT_OUT / "training_history.json").write_text(json.dumps(history, indent=2))
|
|
517
|
+
else:
|
|
518
|
+
patience_counter += 1
|
|
519
|
+
if patience_counter >= PATIENCE:
|
|
520
|
+
console.print(f"\n [bold yellow]Early stopping at epoch {epoch}[/bold yellow] — val_loss stagnant for {PATIENCE} epochs.")
|
|
521
|
+
break
|
|
522
|
+
|
|
523
|
+
console.print(scorecard)
|
|
524
|
+
console.print(f"\n[bold green]SecureBERT Complete[/bold green]: Best val_loss={best_val_loss:.4f} Best micro_F1={best_f1:.3f}")
|
|
525
|
+
console.print(f" Checkpoint saved to: [white]{SECUREBERT_OUT}[/white]\n")
|
|
526
|
+
return {"best_val_loss": best_val_loss, "best_micro_f1": best_f1, "epochs": len(history), "history": history}
|
|
527
|
+
|
|
528
|
+
|
|
529
|
+
# ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━
|
|
530
|
+
# STEP 3: Laya System 1 Fine-Tuning (Augmented with MoreFixes Route Signals)
|
|
531
|
+
# ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━
|
|
532
|
+
|
|
533
|
+
def mine_route_patterns_from_patches() -> List[Any]:
|
|
534
|
+
"""Mines REST route signatures from MoreFixes patch files to augment Laya training.
|
|
535
|
+
|
|
536
|
+
MoreFixes patches are git diffs of real web frameworks. They contain route
|
|
537
|
+
definitions like:
|
|
538
|
+
@app.route('/api/v1/orders/<id>')
|
|
539
|
+
router.get('/api/users/:id', ...)
|
|
540
|
+
@GetMapping('/admin/reset')
|
|
541
|
+
@PreAuthorize('permitAll()') @PostMapping('/admin/config')
|
|
542
|
+
|
|
543
|
+
We parse these out and convert them into the same APM topology format that
|
|
544
|
+
Laya is trained on (Endpoint/Auth/Params/Database/etc.), enriching the
|
|
545
|
+
router's training signal with REAL production endpoint patterns from 9,972
|
|
546
|
+
repositories across Python, Java, Node.js, Go, PHP, Ruby, and Rust.
|
|
547
|
+
"""
|
|
548
|
+
from trace_engine.intelligence.training.laya_trainer import LayaTrainingSample
|
|
549
|
+
|
|
550
|
+
console.print("\n[bold]Mining route patterns from MoreFixes patches for Laya augmentation...[/bold]")
|
|
551
|
+
|
|
552
|
+
# Route detection regex patterns (framework-agnostic)
|
|
553
|
+
ROUTE_PATTERNS = [
|
|
554
|
+
# Flask / FastAPI
|
|
555
|
+
(re.compile(r'@app\.route\(["\']([^"\']+)["\'],\s*methods=\[([^\]]+)\]', re.I), "python"),
|
|
556
|
+
(re.compile(r'@app\.(get|post|put|delete|patch)\(["\']([^"\']+)["\']', re.I), "python"),
|
|
557
|
+
(re.compile(r'@router\.(get|post|put|delete|patch)\(["\']([^"\']+)["\']', re.I), "python"),
|
|
558
|
+
# Express/Node
|
|
559
|
+
(re.compile(r'router\.(get|post|put|delete|patch)\(["\']([^"\']+)["\']', re.I), "node"),
|
|
560
|
+
(re.compile(r'app\.(get|post|put|delete)\(["\']([^"\']+)["\']', re.I), "node"),
|
|
561
|
+
# Spring
|
|
562
|
+
(re.compile(r'@(GetMapping|PostMapping|PutMapping|DeleteMapping)\(["\']([^"\']+)["\']', re.I), "spring"),
|
|
563
|
+
(re.compile(r'@RequestMapping\(value\s*=\s*["\']([^"\']+)["\'],\s*method\s*=\s*RequestMethod\.(\w+)', re.I), "spring"),
|
|
564
|
+
# Django
|
|
565
|
+
(re.compile(r'path\(["\']([^"\']+)["\'],\s*\w+', re.I), "django"),
|
|
566
|
+
# Generic
|
|
567
|
+
(re.compile(r'(GET|POST|PUT|DELETE|PATCH)\s+(/api/[^\s"\']+)', re.I), "generic"),
|
|
568
|
+
]
|
|
569
|
+
|
|
570
|
+
# Auth indicators in patch context
|
|
571
|
+
AUTH_PATTERNS = re.compile(
|
|
572
|
+
r'@login_required|@require_auth|@authenticated|@PreAuthorize|requiresAuth|'
|
|
573
|
+
r'Authorization.*Bearer|session\.get.*user|request\.user\b|getCurrentUser|'
|
|
574
|
+
r'verify_token|check_permission', re.I
|
|
575
|
+
)
|
|
576
|
+
# Database sink indicators
|
|
577
|
+
DB_PATTERNS = re.compile(
|
|
578
|
+
r'\.execute\(|\.query\(|\.findById|\.findOne|db\.|cursor\.|repository\.|'
|
|
579
|
+
r'\.save\(|\.update\(|\.delete\(|SELECT |INSERT |UPDATE |DELETE ', re.I
|
|
580
|
+
)
|
|
581
|
+
# Outbound HTTP indicators
|
|
582
|
+
NET_PATTERNS = re.compile(r'requests\.(get|post)|httpx\.|axios\.|fetch\(|http\.request\(', re.I)
|
|
583
|
+
# Admin/privileged path indicator
|
|
584
|
+
ADMIN_PATH = re.compile(r'/admin/|/manage/|/superadmin/|/ops/|/internal/', re.I)
|
|
585
|
+
# Sensitive parameter names
|
|
586
|
+
SENS_PARAMS = re.compile(r'password|token|secret|key|credential|auth', re.I)
|
|
587
|
+
|
|
588
|
+
samples = []
|
|
589
|
+
processed_routes = set()
|
|
590
|
+
|
|
591
|
+
with zipfile.ZipFile(PATCHES_FILE, "r") as zf:
|
|
592
|
+
for name in zf.namelist():
|
|
593
|
+
if not (name.endswith(".patch") or name.endswith(".diff")):
|
|
594
|
+
continue
|
|
595
|
+
try:
|
|
596
|
+
diff_text = zf.read(name).decode("utf-8", errors="replace")
|
|
597
|
+
except Exception:
|
|
598
|
+
continue
|
|
599
|
+
|
|
600
|
+
# Look for route definitions in ADDED lines (the fixed/new code)
|
|
601
|
+
added_lines = "\n".join(
|
|
602
|
+
line[1:] for line in diff_text.splitlines()
|
|
603
|
+
if line.startswith("+") and not line.startswith("+++")
|
|
604
|
+
)
|
|
605
|
+
removed_lines = "\n".join(
|
|
606
|
+
line[1:] for line in diff_text.splitlines()
|
|
607
|
+
if line.startswith("-") and not line.startswith("---")
|
|
608
|
+
)
|
|
609
|
+
|
|
610
|
+
for pat, framework in ROUTE_PATTERNS:
|
|
611
|
+
for match in pat.finditer(added_lines):
|
|
612
|
+
groups = match.groups()
|
|
613
|
+
# Extract verb and path from match groups (varies by pattern)
|
|
614
|
+
if framework == "spring":
|
|
615
|
+
spring_verb_map = {
|
|
616
|
+
"GetMapping": "GET", "PostMapping": "POST",
|
|
617
|
+
"PutMapping": "PUT", "DeleteMapping": "DELETE",
|
|
618
|
+
}
|
|
619
|
+
verb = spring_verb_map.get(groups[0], "POST")
|
|
620
|
+
path = groups[1]
|
|
621
|
+
elif framework == "python" and "methods" in pat.pattern:
|
|
622
|
+
path = groups[0]
|
|
623
|
+
raw_verb = groups[1].replace('"', "").replace("'", "").split(",")[0].strip()
|
|
624
|
+
verb = raw_verb.upper()
|
|
625
|
+
elif framework == "generic":
|
|
626
|
+
verb = groups[0].upper()
|
|
627
|
+
path = groups[1]
|
|
628
|
+
else:
|
|
629
|
+
verb = groups[0].upper() if groups[0].upper() in ("GET","POST","PUT","DELETE","PATCH") else "GET"
|
|
630
|
+
path = groups[1] if len(groups) > 1 else groups[0]
|
|
631
|
+
|
|
632
|
+
# Deduplicate
|
|
633
|
+
route_key = f"{verb}:{path}"
|
|
634
|
+
if route_key in processed_routes:
|
|
635
|
+
continue
|
|
636
|
+
processed_routes.add(route_key)
|
|
637
|
+
|
|
638
|
+
# Extract features
|
|
639
|
+
has_auth = bool(AUTH_PATTERNS.search(added_lines))
|
|
640
|
+
has_db = bool(DB_PATTERNS.search(added_lines))
|
|
641
|
+
has_net = bool(NET_PATTERNS.search(added_lines))
|
|
642
|
+
is_admin = bool(ADMIN_PATH.search(path))
|
|
643
|
+
is_state_changing = verb in ("POST", "PUT", "DELETE", "PATCH")
|
|
644
|
+
|
|
645
|
+
# Extract parameter names from path template and query params
|
|
646
|
+
params = re.findall(r'[:{<](\w+)[}>]?', path)
|
|
647
|
+
has_sensitive_param = bool(SENS_PARAMS.search(" ".join(params)))
|
|
648
|
+
|
|
649
|
+
# Determine testpack from evidence signals
|
|
650
|
+
if is_admin and not has_auth:
|
|
651
|
+
testpack = "bfla"
|
|
652
|
+
priority = "critical"
|
|
653
|
+
elif has_net and not is_state_changing:
|
|
654
|
+
testpack = "ssrf"
|
|
655
|
+
priority = "critical" if not has_auth else "high"
|
|
656
|
+
elif has_db and params and verb == "GET" and not is_admin:
|
|
657
|
+
testpack = "bola"
|
|
658
|
+
priority = "critical" if not has_auth else "high"
|
|
659
|
+
elif has_db and params and verb in ("POST", "PUT") and not is_admin:
|
|
660
|
+
testpack = "injection"
|
|
661
|
+
priority = "critical"
|
|
662
|
+
elif is_state_changing and not has_auth and has_sensitive_param:
|
|
663
|
+
testpack = "authentication"
|
|
664
|
+
priority = "critical"
|
|
665
|
+
elif is_state_changing and not has_auth:
|
|
666
|
+
testpack = "mass_assignment"
|
|
667
|
+
priority = "high"
|
|
668
|
+
else:
|
|
669
|
+
testpack = "none"
|
|
670
|
+
priority = "low"
|
|
671
|
+
|
|
672
|
+
# Build APM state string in Laya's expected format
|
|
673
|
+
param_list = str(params[:5]) if params else "[]"
|
|
674
|
+
sink_part = "Sinks: 1 detected (DatabaseAccess lookup)" if has_db else (
|
|
675
|
+
"Sinks: 1 detected (OutboundHTTPClient dispatch)" if has_net else
|
|
676
|
+
"No direct sensitive sink"
|
|
677
|
+
)
|
|
678
|
+
admin_roles = '["admin"]'
|
|
679
|
+
roles_str = admin_roles if is_admin else "[]"
|
|
680
|
+
state = (
|
|
681
|
+
f"Endpoint: {verb} {path}\n"
|
|
682
|
+
f"Auth Required: {has_auth}, Roles: {roles_str}\n"
|
|
683
|
+
f"Parameters: {param_list}\n"
|
|
684
|
+
f"Database Access: {has_db}, Outbound Network: {has_net}\n"
|
|
685
|
+
f"State Changing: {is_state_changing}, Sensitive Data: {has_sensitive_param}\n"
|
|
686
|
+
f"APM Path Context: {sink_part}"
|
|
687
|
+
)
|
|
688
|
+
|
|
689
|
+
# Domain group based on path prefix (for potential disjoint analysis)
|
|
690
|
+
domain = "morefixes_mined"
|
|
691
|
+
|
|
692
|
+
samples.append(LayaTrainingSample(
|
|
693
|
+
endpoint_state=state,
|
|
694
|
+
priority_level=priority,
|
|
695
|
+
primary_testpack=testpack,
|
|
696
|
+
domain_group=domain,
|
|
697
|
+
))
|
|
698
|
+
|
|
699
|
+
console.print(f" => Mined [bold green]{len(samples):,}[/bold green] real endpoint patterns from {len(processed_routes):,} unique routes")
|
|
700
|
+
return samples
|
|
701
|
+
|
|
702
|
+
|
|
703
|
+
def train_laya(mined_samples: List[Any]) -> Dict[str, Any]:
|
|
704
|
+
"""Fine-tunes Laya System 1 router (DistilBERT dual-head) on synthetic + mined APM samples.
|
|
705
|
+
|
|
706
|
+
Epoch policy (scientific basis):
|
|
707
|
+
──────────────────────────────────
|
|
708
|
+
DistilBERT-base (66M params) fine-tuned on a SMALL dataset (471 synthetic + mined).
|
|
709
|
+
Mosbach et al. 2020 "On the Stability of Fine-Tuning BERT" found:
|
|
710
|
+
• Small datasets (<1000 samples) are prone to degenerate fine-tuning runs.
|
|
711
|
+
• A longer schedule (20-30 epochs) with a warmup phase dramatically stabilizes
|
|
712
|
+
convergence and achieves better final accuracy than the standard 3-5 epoch recipe.
|
|
713
|
+
• A 10% linear warmup followed by cosine decay is the recommended schedule.
|
|
714
|
+
|
|
715
|
+
We run MAX_EPOCHS=25 with patience=4. The disjoint domain validation split
|
|
716
|
+
(healthcare, fintech, iot, webhooks, admin_tenants) is never seen during training,
|
|
717
|
+
guaranteeing the val_loss reflects TRUE generalization, not memorization.
|
|
718
|
+
|
|
719
|
+
The mined samples from MoreFixes are tagged domain='morefixes_mined' and are
|
|
720
|
+
used for training only (not validation), giving Laya real production URL patterns.
|
|
721
|
+
"""
|
|
722
|
+
MAX_EPOCHS = 25
|
|
723
|
+
PATIENCE = 4
|
|
724
|
+
BATCH_SIZE = 16
|
|
725
|
+
LR = 3e-5
|
|
726
|
+
WEIGHT_DECAY = 0.01
|
|
727
|
+
MODEL_NAME = "distilbert/distilbert-base-uncased"
|
|
728
|
+
WARMUP_RATIO = 0.10
|
|
729
|
+
|
|
730
|
+
console.rule("[bold green]Laya System 1 Fine-Tuning[/bold green]")
|
|
731
|
+
console.print(f"""
|
|
732
|
+
Architecture : DistilBERT-base (66M params) → dual-head (Priority ×4, Testpack ×7)
|
|
733
|
+
Epoch policy : Max {MAX_EPOCHS}, early stopping patience={PATIENCE} on val_loss
|
|
734
|
+
Schedule : 10% linear warmup → cosine decay (Mosbach 2020)
|
|
735
|
+
Optimizer : AdamW lr={LR}, weight_decay={WEIGHT_DECAY}
|
|
736
|
+
Loss : CrossEntropyLoss (multi-task: priority + testpack)
|
|
737
|
+
Val strategy : Disjoint domain split — healthcare/fintech/iot/webhooks/admin_tenants
|
|
738
|
+
MoreFixes : {len(mined_samples):,} real route patterns added to TRAIN only
|
|
739
|
+
""")
|
|
740
|
+
|
|
741
|
+
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
|
742
|
+
console.print(f" GPU: [bold green]{torch.cuda.get_device_name(0) if device.type == 'cuda' else 'CPU'}[/bold green]\n")
|
|
743
|
+
|
|
744
|
+
# ── Dataset assembly ──
|
|
745
|
+
synthetic_corpus = generate_laya_training_corpus()
|
|
746
|
+
all_samples = synthetic_corpus + mined_samples
|
|
747
|
+
|
|
748
|
+
# Disjoint split: held-out domains stay in val, mined samples go to train
|
|
749
|
+
val_samples = [s for s in synthetic_corpus if s.domain_group in VAL_DOMAINS]
|
|
750
|
+
train_samples = [s for s in synthetic_corpus if s.domain_group not in VAL_DOMAINS] + mined_samples
|
|
751
|
+
|
|
752
|
+
console.print(
|
|
753
|
+
f" Synthetic corpus : {len(synthetic_corpus):,} samples\n"
|
|
754
|
+
f" MoreFixes mined : [cyan]{len(mined_samples):,}[/cyan] real production routes\n"
|
|
755
|
+
f" Train set : [bold white]{len(train_samples):,}[/bold white] (synthetic train + all mined)\n"
|
|
756
|
+
f" Val set (disjoint) : [bold green]{len(val_samples):,}[/bold green] samples (held-out domains)\n"
|
|
757
|
+
)
|
|
758
|
+
|
|
759
|
+
from collections import Counter as C
|
|
760
|
+
tc = C(s.primary_testpack for s in train_samples)
|
|
761
|
+
vc = C(s.primary_testpack for s in val_samples)
|
|
762
|
+
console.print(f" Train distribution : {dict(tc)}")
|
|
763
|
+
console.print(f" Val distribution : {dict(vc)}\n")
|
|
764
|
+
|
|
765
|
+
tokenizer = AutoTokenizer.from_pretrained(MODEL_NAME)
|
|
766
|
+
train_ds = LayaDataset(train_samples, tokenizer)
|
|
767
|
+
val_ds = LayaDataset(val_samples, tokenizer)
|
|
768
|
+
train_loader = DataLoader(train_ds, batch_size=BATCH_SIZE, shuffle=True, num_workers=0)
|
|
769
|
+
val_loader = DataLoader(val_ds, batch_size=BATCH_SIZE, shuffle=False, num_workers=0)
|
|
770
|
+
|
|
771
|
+
model = LayaDualHeadModel(base_model_name=MODEL_NAME).to(device)
|
|
772
|
+
optimizer = torch.optim.AdamW(model.parameters(), lr=LR, weight_decay=WEIGHT_DECAY)
|
|
773
|
+
total_steps = len(train_loader) * MAX_EPOCHS
|
|
774
|
+
warmup_steps = int(total_steps * WARMUP_RATIO)
|
|
775
|
+
scheduler = get_linear_schedule_with_warmup(optimizer, warmup_steps, total_steps)
|
|
776
|
+
criterion = nn.CrossEntropyLoss()
|
|
777
|
+
|
|
778
|
+
# ── Training ──
|
|
779
|
+
best_val_loss = float("inf")
|
|
780
|
+
patience_counter = 0
|
|
781
|
+
history = []
|
|
782
|
+
|
|
783
|
+
scorecard = Table(title="[bold green]Laya System 1 Honest Scorecard (Held-Out Domains)[/bold green]", show_header=True)
|
|
784
|
+
scorecard.add_column("Ep", justify="center", style="bold cyan")
|
|
785
|
+
scorecard.add_column("Train Loss", justify="right")
|
|
786
|
+
scorecard.add_column("Val Loss", justify="right", style="yellow")
|
|
787
|
+
scorecard.add_column("Priority Acc", justify="right")
|
|
788
|
+
scorecard.add_column("Testpack Acc", justify="right", style="bold green")
|
|
789
|
+
scorecard.add_column("Macro P", justify="right")
|
|
790
|
+
scorecard.add_column("Macro R", justify="right")
|
|
791
|
+
scorecard.add_column("Macro F1", justify="right", style="bold green")
|
|
792
|
+
scorecard.add_column("Status", justify="center")
|
|
793
|
+
|
|
794
|
+
console.print(f"[bold green]Starting training (max {MAX_EPOCHS} epochs, patience={PATIENCE})...[/bold green]\n")
|
|
795
|
+
|
|
796
|
+
for epoch in range(1, MAX_EPOCHS + 1):
|
|
797
|
+
t0 = time.perf_counter()
|
|
798
|
+
model.train()
|
|
799
|
+
train_loss_sum = 0.0
|
|
800
|
+
|
|
801
|
+
for batch in train_loader:
|
|
802
|
+
input_ids = batch["input_ids"].to(device)
|
|
803
|
+
mask = batch["attention_mask"].to(device)
|
|
804
|
+
prio_lbl = batch["priority_label"].to(device)
|
|
805
|
+
pack_lbl = batch["testpack_label"].to(device)
|
|
806
|
+
|
|
807
|
+
optimizer.zero_grad()
|
|
808
|
+
prio_logits, pack_logits = model(input_ids, mask)
|
|
809
|
+
loss = criterion(prio_logits, prio_lbl) + criterion(pack_logits, pack_lbl)
|
|
810
|
+
loss.backward()
|
|
811
|
+
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
|
|
812
|
+
optimizer.step()
|
|
813
|
+
scheduler.step()
|
|
814
|
+
train_loss_sum += loss.item()
|
|
815
|
+
|
|
816
|
+
avg_train_loss = round(train_loss_sum / len(train_loader), 4)
|
|
817
|
+
|
|
818
|
+
# Validation
|
|
819
|
+
model.eval()
|
|
820
|
+
val_loss_sum = 0.0
|
|
821
|
+
prio_preds_all, prio_tgts_all = [], []
|
|
822
|
+
pack_preds_all, pack_tgts_all = [], []
|
|
823
|
+
|
|
824
|
+
with torch.no_grad():
|
|
825
|
+
for batch in val_loader:
|
|
826
|
+
input_ids = batch["input_ids"].to(device)
|
|
827
|
+
mask = batch["attention_mask"].to(device)
|
|
828
|
+
prio_lbl = batch["priority_label"].to(device)
|
|
829
|
+
pack_lbl = batch["testpack_label"].to(device)
|
|
830
|
+
prio_logits, pack_logits = model(input_ids, mask)
|
|
831
|
+
loss = criterion(prio_logits, prio_lbl) + criterion(pack_logits, pack_lbl)
|
|
832
|
+
val_loss_sum += loss.item()
|
|
833
|
+
prio_preds_all.extend(torch.argmax(prio_logits, dim=-1).cpu().tolist())
|
|
834
|
+
prio_tgts_all.extend(prio_lbl.cpu().tolist())
|
|
835
|
+
pack_preds_all.extend(torch.argmax(pack_logits, dim=-1).cpu().tolist())
|
|
836
|
+
pack_tgts_all.extend(pack_lbl.cpu().tolist())
|
|
837
|
+
|
|
838
|
+
avg_val_loss = round(val_loss_sum / len(val_loader), 4)
|
|
839
|
+
prio_acc = round(100 * sum(p == t for p, t in zip(prio_preds_all, prio_tgts_all)) / len(prio_tgts_all), 1)
|
|
840
|
+
pack_acc = round(100 * sum(p == t for p, t in zip(pack_preds_all, pack_tgts_all)) / len(pack_tgts_all), 1)
|
|
841
|
+
|
|
842
|
+
# Macro F1 for testpack (active classes only)
|
|
843
|
+
prec_s, rec_s, f1_s = [], [], []
|
|
844
|
+
for c_idx, c_name in enumerate(TESTPACK_LABELS):
|
|
845
|
+
if not any(t == c_idx for t in pack_tgts_all):
|
|
846
|
+
continue
|
|
847
|
+
tp = sum(1 for p, t in zip(pack_preds_all, pack_tgts_all) if p == c_idx and t == c_idx)
|
|
848
|
+
fp = sum(1 for p, t in zip(pack_preds_all, pack_tgts_all) if p == c_idx and t != c_idx)
|
|
849
|
+
fn = sum(1 for p, t in zip(pack_preds_all, pack_tgts_all) if p != c_idx and t == c_idx)
|
|
850
|
+
prec = tp / (tp + fp) if (tp + fp) > 0 else 0.0
|
|
851
|
+
rec = tp / (tp + fn) if (tp + fn) > 0 else 0.0
|
|
852
|
+
f1 = 2 * prec * rec / (prec + rec) if (prec + rec) > 0 else 0.0
|
|
853
|
+
prec_s.append(prec); rec_s.append(rec); f1_s.append(f1)
|
|
854
|
+
macro_prec = round((sum(prec_s) / len(prec_s)) * 100, 1) if prec_s else 0.0
|
|
855
|
+
macro_rec = round((sum(rec_s) / len(rec_s)) * 100, 1) if rec_s else 0.0
|
|
856
|
+
macro_f1 = round(sum(f1_s) / len(f1_s), 3) if f1_s else 0.0
|
|
857
|
+
|
|
858
|
+
elapsed = round(time.perf_counter() - t0, 1)
|
|
859
|
+
improved = avg_val_loss < best_val_loss
|
|
860
|
+
|
|
861
|
+
scorecard.add_row(
|
|
862
|
+
str(epoch), str(avg_train_loss), str(avg_val_loss),
|
|
863
|
+
f"{prio_acc}%", f"{pack_acc}%",
|
|
864
|
+
f"{macro_prec}%", f"{macro_rec}%", str(macro_f1),
|
|
865
|
+
"[green]BEST[/green]" if improved else f"Stagnant {patience_counter + 1}/{PATIENCE}",
|
|
866
|
+
)
|
|
867
|
+
console.print(
|
|
868
|
+
f" Epoch {epoch:>2}/{MAX_EPOCHS}: train={avg_train_loss} val={avg_val_loss} "
|
|
869
|
+
f"priority_acc={prio_acc}% testpack_acc={pack_acc}% "
|
|
870
|
+
f"macro_P={macro_prec}% macro_R={macro_rec}% macro_F1={macro_f1} [{elapsed}s]"
|
|
871
|
+
)
|
|
872
|
+
|
|
873
|
+
history.append({
|
|
874
|
+
"epoch": epoch, "train_loss": avg_train_loss, "val_loss": avg_val_loss,
|
|
875
|
+
"priority_acc": prio_acc, "testpack_acc": pack_acc,
|
|
876
|
+
"macro_prec": macro_prec, "macro_rec": macro_rec, "macro_f1": macro_f1,
|
|
877
|
+
})
|
|
878
|
+
|
|
879
|
+
if improved:
|
|
880
|
+
best_val_loss = avg_val_loss
|
|
881
|
+
patience_counter = 0
|
|
882
|
+
LAYA_OUT.mkdir(parents=True, exist_ok=True)
|
|
883
|
+
torch.save(model.state_dict(), LAYA_OUT / "laya_dual_head.pt")
|
|
884
|
+
tokenizer.save_pretrained(LAYA_OUT)
|
|
885
|
+
onnx_path = export_to_onnx(model, tokenizer, LAYA_OUT)
|
|
886
|
+
(LAYA_OUT / "laya_metadata.json").write_text(json.dumps({
|
|
887
|
+
"base_model": MODEL_NAME,
|
|
888
|
+
"priority_labels": PRIORITY_LABELS,
|
|
889
|
+
"testpack_labels": TESTPACK_LABELS,
|
|
890
|
+
"best_val_loss": best_val_loss,
|
|
891
|
+
"priority_acc": prio_acc,
|
|
892
|
+
"testpack_acc": pack_acc,
|
|
893
|
+
"macro_f1": macro_f1,
|
|
894
|
+
"macro_prec": macro_prec,
|
|
895
|
+
"macro_rec": macro_rec,
|
|
896
|
+
"val_domains": sorted(VAL_DOMAINS),
|
|
897
|
+
"epochs_trained": epoch,
|
|
898
|
+
"onnx_available": onnx_path is not None,
|
|
899
|
+
"train_samples": len(train_samples),
|
|
900
|
+
"val_samples": len(val_samples),
|
|
901
|
+
"morefixes_mined": len(mined_samples),
|
|
902
|
+
}, indent=2))
|
|
903
|
+
(LAYA_OUT / "training_history.json").write_text(json.dumps(history, indent=2))
|
|
904
|
+
else:
|
|
905
|
+
patience_counter += 1
|
|
906
|
+
if patience_counter >= PATIENCE:
|
|
907
|
+
console.print(f"\n [bold yellow]Early stopping at epoch {epoch}[/bold yellow] — val_loss stagnant for {PATIENCE} epochs.")
|
|
908
|
+
break
|
|
909
|
+
|
|
910
|
+
console.print(scorecard)
|
|
911
|
+
console.print(f"\n[bold green]Laya System 1 Complete[/bold green]: Best val_loss={best_val_loss:.4f} Epochs trained={len(history)}")
|
|
912
|
+
console.print(f" Checkpoint + ONNX saved to: [white]{LAYA_OUT}[/white]\n")
|
|
913
|
+
return {"best_val_loss": best_val_loss, "epochs": len(history), "history": history}
|
|
914
|
+
|
|
915
|
+
|
|
916
|
+
# ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━
|
|
917
|
+
# STEP 4: Dataset Cleanup
|
|
918
|
+
# ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━
|
|
919
|
+
|
|
920
|
+
def cleanup() -> None:
|
|
921
|
+
console.print("\n[bold red]Deleting raw dataset files from C:\\dataset\\ ...[/bold red]")
|
|
922
|
+
for f in [PATCHES_FILE, SQL_FILE]:
|
|
923
|
+
if f.exists():
|
|
924
|
+
f.unlink()
|
|
925
|
+
console.print(f" Deleted: {f}")
|
|
926
|
+
# Also delete the directory if empty
|
|
927
|
+
try:
|
|
928
|
+
if DATASET_DIR.exists() and not list(DATASET_DIR.iterdir()):
|
|
929
|
+
DATASET_DIR.rmdir()
|
|
930
|
+
console.print(f" Removed empty directory: {DATASET_DIR}")
|
|
931
|
+
except Exception:
|
|
932
|
+
pass
|
|
933
|
+
console.print("[bold green]Cleanup complete.[/bold green]")
|
|
934
|
+
|
|
935
|
+
|
|
936
|
+
# ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━
|
|
937
|
+
# MAIN ORCHESTRATOR
|
|
938
|
+
# ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━
|
|
939
|
+
|
|
940
|
+
def main() -> None:
|
|
941
|
+
t_total = time.perf_counter()
|
|
942
|
+
|
|
943
|
+
console.rule("[bold white]TRACE Dual-Model Training Orchestrator[/bold white]")
|
|
944
|
+
console.print("""
|
|
945
|
+
Training two intelligence engines from MoreFixes (Zenodo 20776007):
|
|
946
|
+
|
|
947
|
+
[bold cyan]Engine 1: SecureBERT[/bold cyan] (Code vulnerability classifier)
|
|
948
|
+
Base model : ehsanaghaei/SecureBERT (RoBERTa-base, 125M params)
|
|
949
|
+
Max epochs : 12 • Patience : 3 (Devlin et al. 2019 + Sun et al. 2020)
|
|
950
|
+
LR : 2e-5 • FP16 AMP • BCEWithLogitsLoss + sqrt pos_weight
|
|
951
|
+
|
|
952
|
+
[bold cyan]Engine 2: Laya System 1[/bold cyan] (REST endpoint testpack router)
|
|
953
|
+
Base model : distilbert-base-uncased (66M params)
|
|
954
|
+
Max epochs : 25 • Patience : 4 (Mosbach et al. 2020: small data)
|
|
955
|
+
LR : 3e-5 • 10% linear warmup + cosine • Grad clip=1.0
|
|
956
|
+
Val split : Strict disjoint domain (healthcare / fintech / iot / webhooks)
|
|
957
|
+
""")
|
|
958
|
+
|
|
959
|
+
# Verify files
|
|
960
|
+
if not PATCHES_FILE.exists() or not SQL_FILE.exists():
|
|
961
|
+
console.print("[bold red]ERROR: Dataset files missing from C:\\dataset\\[/bold red]")
|
|
962
|
+
console.print(f" Expected: {PATCHES_FILE} ({PATCHES_FILE.exists()})")
|
|
963
|
+
console.print(f" Expected: {SQL_FILE} ({SQL_FILE.exists()})")
|
|
964
|
+
sys.exit(1)
|
|
965
|
+
|
|
966
|
+
console.print(f" Patch archive : {PATCHES_FILE.stat().st_size / 1024**3:.2f} GB [green]OK[/green]")
|
|
967
|
+
console.print(f" SQL dump : {SQL_FILE.stat().st_size / 1024**3:.2f} GB [green]OK[/green]\n")
|
|
968
|
+
|
|
969
|
+
# ── Stage 1: Extract MoreFixes samples ──────────────────────────────────────
|
|
970
|
+
console.rule("[bold]Stage 1 of 4 — Parsing + Extracting MoreFixes Data[/bold]")
|
|
971
|
+
hash_to_cwe = parse_cwe_mappings()
|
|
972
|
+
morefixes_samples = extract_patch_samples(hash_to_cwe)
|
|
973
|
+
|
|
974
|
+
# ── Stage 2: Mine route patterns for Laya ───────────────────────────────────
|
|
975
|
+
console.rule("[bold]Stage 2 of 4 — Mining Route Patterns for Laya[/bold]")
|
|
976
|
+
laya_mined = mine_route_patterns_from_patches()
|
|
977
|
+
|
|
978
|
+
# ── Stage 3: SecureBERT training ────────────────────────────────────────────
|
|
979
|
+
console.rule("[bold]Stage 3 of 4 — SecureBERT Fine-Tuning (GPU)[/bold]")
|
|
980
|
+
sb_results = train_securebert(morefixes_samples)
|
|
981
|
+
|
|
982
|
+
# ── Stage 4: Laya training ───────────────────────────────────────────────────
|
|
983
|
+
console.rule("[bold]Stage 4 of 4 — Laya System 1 Fine-Tuning (GPU)[/bold]")
|
|
984
|
+
laya_results = train_laya(laya_mined)
|
|
985
|
+
|
|
986
|
+
# ── Cleanup ─────────────────────────────────────────────────────────────────
|
|
987
|
+
cleanup()
|
|
988
|
+
|
|
989
|
+
# ── Final Summary ────────────────────────────────────────────────────────────
|
|
990
|
+
total_time = round(time.perf_counter() - t_total, 1)
|
|
991
|
+
console.rule("[bold green]Training Complete[/bold green]")
|
|
992
|
+
console.print(f"""
|
|
993
|
+
[bold cyan]SecureBERT[/bold cyan]
|
|
994
|
+
Epochs trained : {sb_results['epochs']}
|
|
995
|
+
Best val_loss : {sb_results['best_val_loss']:.4f}
|
|
996
|
+
Best micro_F1 : {sb_results['best_micro_f1']:.3f}
|
|
997
|
+
Checkpoint : {SECUREBERT_OUT}
|
|
998
|
+
|
|
999
|
+
[bold cyan]Laya System 1[/bold cyan]
|
|
1000
|
+
Epochs trained : {laya_results['epochs']}
|
|
1001
|
+
Best val_loss : {laya_results['best_val_loss']:.4f}
|
|
1002
|
+
Checkpoint + ONNX: {LAYA_OUT}
|
|
1003
|
+
|
|
1004
|
+
Total wall time : {total_time}s
|
|
1005
|
+
""")
|
|
1006
|
+
|
|
1007
|
+
|
|
1008
|
+
if __name__ == "__main__":
|
|
1009
|
+
main()
|