@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,321 @@
|
|
|
1
|
+
"""PyTorch CUDA accelerated model fine-tuning engine for SecureBERT with calibrated multi-label metrics."""
|
|
2
|
+
|
|
3
|
+
import os
|
|
4
|
+
import time
|
|
5
|
+
import json
|
|
6
|
+
import logging
|
|
7
|
+
from pathlib import Path
|
|
8
|
+
from typing import Optional, Dict, Any, List
|
|
9
|
+
from pydantic import BaseModel
|
|
10
|
+
from rich.console import Console
|
|
11
|
+
from rich.table import Table
|
|
12
|
+
from rich.progress import Progress, SpinnerColumn, TextColumn, BarColumn, TimeRemainingColumn
|
|
13
|
+
|
|
14
|
+
import torch
|
|
15
|
+
import torch.nn as nn
|
|
16
|
+
from torch.utils.data import DataLoader, random_split
|
|
17
|
+
from transformers import AutoTokenizer, AutoModelForSequenceClassification
|
|
18
|
+
|
|
19
|
+
from trace_engine.intelligence.securebert.classifier import VULN_CATEGORIES
|
|
20
|
+
from trace_engine.intelligence.training.dataset import (
|
|
21
|
+
VulnerabilityDataset,
|
|
22
|
+
generate_cybersecurity_training_corpus,
|
|
23
|
+
)
|
|
24
|
+
|
|
25
|
+
logger = logging.getLogger(__name__)
|
|
26
|
+
console = Console()
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
class TrainingConfig(BaseModel):
|
|
30
|
+
"""Configuration hyperparameters for model fine-tuning."""
|
|
31
|
+
model_name: str = "ehsanaghaei/SecureBERT"
|
|
32
|
+
epochs: int = 5
|
|
33
|
+
batch_size: int = 32
|
|
34
|
+
learning_rate: float = 3e-5
|
|
35
|
+
weight_decay: float = 0.01
|
|
36
|
+
max_length: int = 256
|
|
37
|
+
val_split: float = 0.15
|
|
38
|
+
fp16: bool = True
|
|
39
|
+
eval_threshold: float = 0.50
|
|
40
|
+
early_stopping: bool = True
|
|
41
|
+
patience: int = 2
|
|
42
|
+
min_delta: float = 0.001
|
|
43
|
+
output_dir: str = ".trace/models/securebert-finetuned"
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
class SecureBERTTrainer:
|
|
47
|
+
"""Trains and fine-tunes SecureBERT using PyTorch with NVIDIA GPU acceleration and calibrated multi-label metrics."""
|
|
48
|
+
|
|
49
|
+
def __init__(self, config: Optional[TrainingConfig] = None):
|
|
50
|
+
self.config = config or TrainingConfig()
|
|
51
|
+
self.device = self._detect_device()
|
|
52
|
+
self.scaler = torch.amp.GradScaler("cuda") if self.device.type == "cuda" and self.config.fp16 else None
|
|
53
|
+
|
|
54
|
+
def _detect_device(self) -> torch.device:
|
|
55
|
+
if torch.cuda.is_available():
|
|
56
|
+
dev_name = torch.cuda.get_device_name(0)
|
|
57
|
+
vram_gb = round(torch.cuda.get_device_properties(0).total_memory / (1024**3), 2)
|
|
58
|
+
console.print(f"\n[bold green]GPU Acceleration Active[/bold green]: {dev_name} ({vram_gb} GB VRAM, Mixed Precision: {'FP16' if self.config.fp16 else 'FP32'})")
|
|
59
|
+
return torch.device("cuda:0")
|
|
60
|
+
else:
|
|
61
|
+
console.print("\n[yellow]CUDA not available. Using CPU.[/yellow]")
|
|
62
|
+
return torch.device("cpu")
|
|
63
|
+
|
|
64
|
+
def train(self) -> Dict[str, Any]:
|
|
65
|
+
"""Executes full supervised fine-tuning pipeline on comprehensive cybersecurity corpus."""
|
|
66
|
+
out_path = Path(self.config.output_dir)
|
|
67
|
+
out_path.mkdir(parents=True, exist_ok=True)
|
|
68
|
+
|
|
69
|
+
console.print(f"[bold]Loading Base Checkpoint & Tokenizer[/bold]: {self.config.model_name}...")
|
|
70
|
+
tokenizer = AutoTokenizer.from_pretrained(self.config.model_name)
|
|
71
|
+
model = AutoModelForSequenceClassification.from_pretrained(
|
|
72
|
+
self.config.model_name,
|
|
73
|
+
num_labels=len(VULN_CATEGORIES),
|
|
74
|
+
ignore_mismatched_sizes=True,
|
|
75
|
+
).to(self.device)
|
|
76
|
+
|
|
77
|
+
# 1. Dataset Generation from Big-Vul, CVEfixes, OWASP Benchmark, parquet CVEs, and API templates
|
|
78
|
+
raw_samples = generate_cybersecurity_training_corpus(multiplier=4, include_external_cve=True)
|
|
79
|
+
full_dataset = VulnerabilityDataset(raw_samples, tokenizer, max_length=self.config.max_length)
|
|
80
|
+
|
|
81
|
+
# Compute dampened positive class weights to balance sparse categories without over-firing
|
|
82
|
+
pos_weights = full_dataset.calculate_pos_weights().to(self.device)
|
|
83
|
+
|
|
84
|
+
val_size = int(len(full_dataset) * self.config.val_split)
|
|
85
|
+
train_size = len(full_dataset) - val_size
|
|
86
|
+
train_data, val_data = random_split(
|
|
87
|
+
full_dataset,
|
|
88
|
+
[train_size, val_size],
|
|
89
|
+
generator=torch.Generator().manual_seed(42),
|
|
90
|
+
)
|
|
91
|
+
|
|
92
|
+
train_loader = DataLoader(train_data, batch_size=self.config.batch_size, shuffle=True)
|
|
93
|
+
val_loader = DataLoader(val_data, batch_size=self.config.batch_size, shuffle=False)
|
|
94
|
+
|
|
95
|
+
vuln_count = sum(1 for _, l in raw_samples if l)
|
|
96
|
+
safe_count = len(raw_samples) - vuln_count
|
|
97
|
+
console.print(
|
|
98
|
+
f"Dataset assembled: [bold green]{len(full_dataset)} total code slices[/bold green] "
|
|
99
|
+
f"({vuln_count} vulnerable, {safe_count} safe controls, {train_size} train, {val_size} val)"
|
|
100
|
+
)
|
|
101
|
+
|
|
102
|
+
# 2. Optimization setup: AdamW with Cosine Annealing & Dampened Weighted BCE Loss
|
|
103
|
+
optimizer = torch.optim.AdamW(
|
|
104
|
+
model.parameters(),
|
|
105
|
+
lr=self.config.learning_rate,
|
|
106
|
+
weight_decay=self.config.weight_decay,
|
|
107
|
+
)
|
|
108
|
+
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=self.config.epochs, eta_min=5e-6)
|
|
109
|
+
criterion = nn.BCEWithLogitsLoss(pos_weight=pos_weights)
|
|
110
|
+
|
|
111
|
+
total_start = time.perf_counter()
|
|
112
|
+
history: List[Dict[str, float]] = []
|
|
113
|
+
best_val_loss = float("inf")
|
|
114
|
+
best_f1 = 0.0
|
|
115
|
+
patience_counter = 0
|
|
116
|
+
standard_th = self.config.eval_threshold
|
|
117
|
+
|
|
118
|
+
console.print(f"\n[bold green]Beginning GPU Fine-Tuning ({self.config.epochs} Epochs on {self.device}, Evaluation Threshold: {standard_th}, Early Stopping: {self.config.early_stopping})[/bold green]:")
|
|
119
|
+
|
|
120
|
+
for epoch in range(1, self.config.epochs + 1):
|
|
121
|
+
epoch_start = time.perf_counter()
|
|
122
|
+
model.train()
|
|
123
|
+
total_train_loss = 0.0
|
|
124
|
+
|
|
125
|
+
with Progress(
|
|
126
|
+
SpinnerColumn(),
|
|
127
|
+
TextColumn(f"[bold cyan]Epoch {epoch}/{self.config.epochs}"),
|
|
128
|
+
BarColumn(),
|
|
129
|
+
TextColumn("[progress.percentage]{task.percentage:>3.0f}%"),
|
|
130
|
+
TimeRemainingColumn(),
|
|
131
|
+
console=console,
|
|
132
|
+
) as progress:
|
|
133
|
+
task = progress.add_task("Training", total=len(train_loader))
|
|
134
|
+
|
|
135
|
+
for batch in train_loader:
|
|
136
|
+
optimizer.zero_grad()
|
|
137
|
+
input_ids = batch["input_ids"].to(self.device)
|
|
138
|
+
attention_mask = batch["attention_mask"].to(self.device)
|
|
139
|
+
labels = batch["labels"].to(self.device)
|
|
140
|
+
|
|
141
|
+
if self.scaler:
|
|
142
|
+
with torch.amp.autocast("cuda"):
|
|
143
|
+
outputs = model(input_ids=input_ids, attention_mask=attention_mask)
|
|
144
|
+
loss = criterion(outputs.logits, labels)
|
|
145
|
+
self.scaler.scale(loss).backward()
|
|
146
|
+
self.scaler.step(optimizer)
|
|
147
|
+
self.scaler.update()
|
|
148
|
+
else:
|
|
149
|
+
outputs = model(input_ids=input_ids, attention_mask=attention_mask)
|
|
150
|
+
loss = criterion(outputs.logits, labels)
|
|
151
|
+
loss.backward()
|
|
152
|
+
optimizer.step()
|
|
153
|
+
|
|
154
|
+
total_train_loss += loss.item()
|
|
155
|
+
progress.update(task, advance=1)
|
|
156
|
+
|
|
157
|
+
scheduler.step()
|
|
158
|
+
avg_train_loss = round(total_train_loss / len(train_loader), 4)
|
|
159
|
+
|
|
160
|
+
# Validation Pass
|
|
161
|
+
model.eval()
|
|
162
|
+
total_val_loss = 0.0
|
|
163
|
+
all_probs: List[torch.Tensor] = []
|
|
164
|
+
all_labels: List[torch.Tensor] = []
|
|
165
|
+
|
|
166
|
+
with torch.no_grad():
|
|
167
|
+
for val_batch in val_loader:
|
|
168
|
+
v_input_ids = val_batch["input_ids"].to(self.device)
|
|
169
|
+
v_mask = val_batch["attention_mask"].to(self.device)
|
|
170
|
+
v_labels = val_batch["labels"].to(self.device)
|
|
171
|
+
|
|
172
|
+
if self.scaler:
|
|
173
|
+
with torch.amp.autocast("cuda"):
|
|
174
|
+
v_outputs = model(input_ids=v_input_ids, attention_mask=v_mask)
|
|
175
|
+
v_loss = criterion(v_outputs.logits, v_labels)
|
|
176
|
+
else:
|
|
177
|
+
v_outputs = model(input_ids=v_input_ids, attention_mask=v_mask)
|
|
178
|
+
v_loss = criterion(v_outputs.logits, v_labels)
|
|
179
|
+
|
|
180
|
+
total_val_loss += v_loss.item()
|
|
181
|
+
probs = torch.sigmoid(v_outputs.logits)
|
|
182
|
+
all_probs.append(probs.cpu())
|
|
183
|
+
all_labels.append(v_labels.cpu())
|
|
184
|
+
|
|
185
|
+
avg_val_loss = round(total_val_loss / len(val_loader), 4)
|
|
186
|
+
cat_probs = torch.cat(all_probs, dim=0)
|
|
187
|
+
cat_labels = torch.cat(all_labels, dim=0)
|
|
188
|
+
|
|
189
|
+
# Compute metrics at standard unskewed threshold (0.50)
|
|
190
|
+
preds = (cat_probs >= standard_th).float()
|
|
191
|
+
tp = ((preds == 1) & (cat_labels == 1)).sum().item()
|
|
192
|
+
fp = ((preds == 1) & (cat_labels == 0)).sum().item()
|
|
193
|
+
fn = ((preds == 0) & (cat_labels == 1)).sum().item()
|
|
194
|
+
|
|
195
|
+
p = round(tp / (tp + fp), 3) if (tp + fp) > 0 else 0.0
|
|
196
|
+
r = round(tp / (tp + fn), 3) if (tp + fn) > 0 else 0.0
|
|
197
|
+
micro_f1 = round((2 * p * r / (p + r)), 3) if (p + r) > 0 else 0.0
|
|
198
|
+
|
|
199
|
+
# Macro-F1 across all active categories in validation set
|
|
200
|
+
per_class_f1: List[float] = []
|
|
201
|
+
for c_idx in range(len(VULN_CATEGORIES)):
|
|
202
|
+
c_pred = preds[:, c_idx]
|
|
203
|
+
c_true = cat_labels[:, c_idx]
|
|
204
|
+
c_tp = ((c_pred == 1) & (c_true == 1)).sum().item()
|
|
205
|
+
c_fp = ((c_pred == 1) & (c_true == 0)).sum().item()
|
|
206
|
+
c_fn = ((c_pred == 0) & (c_true == 1)).sum().item()
|
|
207
|
+
c_p = c_tp / (c_tp + c_fp) if (c_tp + c_fp) > 0 else 0.0
|
|
208
|
+
c_r = c_tp / (c_tp + c_fn) if (c_tp + c_fn) > 0 else 0.0
|
|
209
|
+
c_f1 = (2 * c_p * c_r / (c_p + c_r)) if (c_p + c_r) > 0 else 0.0
|
|
210
|
+
if (c_true == 1).sum().item() > 0:
|
|
211
|
+
per_class_f1.append(c_f1)
|
|
212
|
+
macro_f1 = round(sum(per_class_f1) / len(per_class_f1), 3) if per_class_f1 else 0.0
|
|
213
|
+
|
|
214
|
+
# Hamming accuracy (overall multi-label decision accuracy across all labels)
|
|
215
|
+
hamming_acc = round((preds == cat_labels).float().mean().item(), 4)
|
|
216
|
+
|
|
217
|
+
# Exact-Match (Subset Accuracy): all 10 labels must match simultaneously
|
|
218
|
+
subset_acc = round((preds == cat_labels).all(dim=-1).float().mean().item(), 4)
|
|
219
|
+
|
|
220
|
+
# Top-1 accuracy on vulnerable samples
|
|
221
|
+
vuln_mask = cat_labels.sum(dim=-1) > 0
|
|
222
|
+
if vuln_mask.sum() > 0:
|
|
223
|
+
top1_preds = cat_probs[vuln_mask].argmax(dim=-1)
|
|
224
|
+
top1_labels = cat_labels[vuln_mask].argmax(dim=-1)
|
|
225
|
+
top1_acc = round((top1_preds == top1_labels).float().mean().item(), 4)
|
|
226
|
+
else:
|
|
227
|
+
top1_acc = 1.0
|
|
228
|
+
|
|
229
|
+
duration = round(time.perf_counter() - epoch_start, 2)
|
|
230
|
+
|
|
231
|
+
console.print(
|
|
232
|
+
f" Epoch {epoch}: Train Loss: [cyan]{avg_train_loss}[/cyan] | "
|
|
233
|
+
f"Val Loss: [yellow]{avg_val_loss}[/yellow] | "
|
|
234
|
+
f"Hamming Acc: [bold white]{round(hamming_acc * 100, 1)}%[/bold white] | "
|
|
235
|
+
f"Exact Match: [bold white]{round(subset_acc * 100, 1)}%[/bold white] | "
|
|
236
|
+
f"Top-1: [bold white]{round(top1_acc * 100, 1)}%[/bold white] | "
|
|
237
|
+
f"Precision: [green]{int(p * 100)}%[/green] | "
|
|
238
|
+
f"Recall: [green]{int(r * 100)}%[/green] | "
|
|
239
|
+
f"Micro F1: [bold green]{micro_f1}[/bold green] | "
|
|
240
|
+
f"Macro F1: [bold green]{macro_f1}[/bold green] [{duration}s]"
|
|
241
|
+
)
|
|
242
|
+
|
|
243
|
+
history.append({
|
|
244
|
+
"epoch": epoch,
|
|
245
|
+
"train_loss": avg_train_loss,
|
|
246
|
+
"val_loss": avg_val_loss,
|
|
247
|
+
"hamming_accuracy": hamming_acc,
|
|
248
|
+
"exact_match_accuracy": subset_acc,
|
|
249
|
+
"top1_accuracy": top1_acc,
|
|
250
|
+
"precision": p,
|
|
251
|
+
"recall": r,
|
|
252
|
+
"micro_f1": micro_f1,
|
|
253
|
+
"macro_f1": macro_f1,
|
|
254
|
+
"threshold": standard_th,
|
|
255
|
+
})
|
|
256
|
+
|
|
257
|
+
# Checkpoint save on improved val_loss or peak micro F1
|
|
258
|
+
val_improved = avg_val_loss < (best_val_loss - self.config.min_delta)
|
|
259
|
+
if val_improved or (micro_f1 > (best_f1 + self.config.min_delta)):
|
|
260
|
+
if val_improved:
|
|
261
|
+
best_val_loss = avg_val_loss
|
|
262
|
+
best_f1 = max(best_f1, micro_f1)
|
|
263
|
+
patience_counter = 0
|
|
264
|
+
model.save_pretrained(out_path)
|
|
265
|
+
tokenizer.save_pretrained(out_path)
|
|
266
|
+
else:
|
|
267
|
+
patience_counter += 1
|
|
268
|
+
if self.config.early_stopping and patience_counter >= self.config.patience:
|
|
269
|
+
console.print(
|
|
270
|
+
f"\n[bold yellow][EARLY STOPPING TRIGGERED][/bold yellow] Validation loss did not improve for {patience_counter} consecutive epochs "
|
|
271
|
+
f"(Patience threshold: {self.config.patience}, Best Val Loss: {best_val_loss:.4f}, Optimal Checkpoint F1: {best_f1}). "
|
|
272
|
+
f"Gracefully halting training and generating final scorecard."
|
|
273
|
+
)
|
|
274
|
+
break
|
|
275
|
+
|
|
276
|
+
total_time = round(time.perf_counter() - total_start, 2)
|
|
277
|
+
console.print(f"\n[bold green][SUCCESS] Training Complete in {total_time}s! Peak Micro F1: {best_f1}[/bold green]")
|
|
278
|
+
|
|
279
|
+
# 3. Final Summary Table
|
|
280
|
+
table = Table(title="SecureBERT 2.0 Genuine Fine-Tuning Performance (Big-Vul + CVEfixes)", header_style="bold green")
|
|
281
|
+
table.add_column("Epoch", style="cyan")
|
|
282
|
+
table.add_column("Train Loss", justify="right")
|
|
283
|
+
table.add_column("Val Loss", justify="right")
|
|
284
|
+
table.add_column("Hamming Acc", justify="right")
|
|
285
|
+
table.add_column("Exact Match", justify="right", style="bold white")
|
|
286
|
+
table.add_column("Top-1 Acc", justify="right")
|
|
287
|
+
table.add_column("Precision", justify="right", style="green")
|
|
288
|
+
table.add_column("Recall", justify="right", style="green")
|
|
289
|
+
table.add_column("Micro F1", justify="right", style="bold green")
|
|
290
|
+
table.add_column("Macro F1", justify="right", style="bold cyan")
|
|
291
|
+
|
|
292
|
+
for h in history:
|
|
293
|
+
table.add_row(
|
|
294
|
+
str(h["epoch"]),
|
|
295
|
+
str(h["train_loss"]),
|
|
296
|
+
str(h["val_loss"]),
|
|
297
|
+
f"{round(h['hamming_accuracy'] * 100, 1)}%",
|
|
298
|
+
f"{round(h['exact_match_accuracy'] * 100, 1)}%",
|
|
299
|
+
f"{round(h['top1_accuracy'] * 100, 1)}%",
|
|
300
|
+
f"{int(h['precision'] * 100)}%",
|
|
301
|
+
f"{int(h['recall'] * 100)}%",
|
|
302
|
+
str(h["micro_f1"]),
|
|
303
|
+
str(h["macro_f1"]),
|
|
304
|
+
)
|
|
305
|
+
console.print(table)
|
|
306
|
+
|
|
307
|
+
metadata = {
|
|
308
|
+
"base_model": self.config.model_name,
|
|
309
|
+
"categories": VULN_CATEGORIES,
|
|
310
|
+
"epochs": self.config.epochs,
|
|
311
|
+
"peak_f1": best_f1,
|
|
312
|
+
"eval_threshold": standard_th,
|
|
313
|
+
"device": str(self.device),
|
|
314
|
+
"training_time_seconds": total_time,
|
|
315
|
+
"history": history,
|
|
316
|
+
}
|
|
317
|
+
with open(out_path / "training_metadata.json", "w", encoding="utf-8") as f:
|
|
318
|
+
json.dump(metadata, f, indent=2)
|
|
319
|
+
|
|
320
|
+
console.print(f"[bold green]Best model checkpoint deployed to {out_path}![/bold green]\n")
|
|
321
|
+
return metadata
|