@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.
Files changed (249) hide show
  1. package/README.md +208 -0
  2. package/bin/trace.js +112 -0
  3. package/package.json +46 -0
  4. package/pyproject.toml +39 -0
  5. package/src/trace_engine/__init__.py +8 -0
  6. package/src/trace_engine/__pycache__/__init__.cpython-311.pyc +0 -0
  7. package/src/trace_engine/__pycache__/cli.cpython-311.pyc +0 -0
  8. package/src/trace_engine/__pycache__/doctor.cpython-311.pyc +0 -0
  9. package/src/trace_engine/__pycache__/interactive.cpython-311.pyc +0 -0
  10. package/src/trace_engine/__pycache__/verify.cpython-311.pyc +0 -0
  11. package/src/trace_engine/ai/__init__.py +7 -0
  12. package/src/trace_engine/ai/__pycache__/__init__.cpython-311.pyc +0 -0
  13. package/src/trace_engine/ai/__pycache__/base.cpython-311.pyc +0 -0
  14. package/src/trace_engine/ai/__pycache__/ollama.cpython-311.pyc +0 -0
  15. package/src/trace_engine/ai/__pycache__/planner.cpython-311.pyc +0 -0
  16. package/src/trace_engine/ai/base.py +23 -0
  17. package/src/trace_engine/ai/ollama.py +50 -0
  18. package/src/trace_engine/ai/planner.py +40 -0
  19. package/src/trace_engine/apm/__init__.py +18 -0
  20. package/src/trace_engine/apm/__pycache__/__init__.cpython-311.pyc +0 -0
  21. package/src/trace_engine/apm/__pycache__/builder.cpython-311.pyc +0 -0
  22. package/src/trace_engine/apm/__pycache__/edges.cpython-311.pyc +0 -0
  23. package/src/trace_engine/apm/__pycache__/model.cpython-311.pyc +0 -0
  24. package/src/trace_engine/apm/__pycache__/nodes.cpython-311.pyc +0 -0
  25. package/src/trace_engine/apm/__pycache__/serialization.cpython-311.pyc +0 -0
  26. package/src/trace_engine/apm/builder.py +208 -0
  27. package/src/trace_engine/apm/edges.py +25 -0
  28. package/src/trace_engine/apm/model.py +107 -0
  29. package/src/trace_engine/apm/nodes.py +27 -0
  30. package/src/trace_engine/apm/serialization.py +105 -0
  31. package/src/trace_engine/benchmark/__init__.py +5 -0
  32. package/src/trace_engine/benchmark/__pycache__/__init__.cpython-311.pyc +0 -0
  33. package/src/trace_engine/benchmark/__pycache__/owasp.cpython-311.pyc +0 -0
  34. package/src/trace_engine/benchmark/owasp.py +183 -0
  35. package/src/trace_engine/cli.py +1184 -0
  36. package/src/trace_engine/config/__init__.py +35 -0
  37. package/src/trace_engine/config/__pycache__/__init__.cpython-311.pyc +0 -0
  38. package/src/trace_engine/config/__pycache__/defaults.cpython-311.pyc +0 -0
  39. package/src/trace_engine/config/__pycache__/loader.cpython-311.pyc +0 -0
  40. package/src/trace_engine/config/__pycache__/settings.cpython-311.pyc +0 -0
  41. package/src/trace_engine/config/defaults.py +48 -0
  42. package/src/trace_engine/config/loader.py +64 -0
  43. package/src/trace_engine/config/settings.py +72 -0
  44. package/src/trace_engine/doctor.py +250 -0
  45. package/src/trace_engine/findings/__init__.py +15 -0
  46. package/src/trace_engine/findings/__pycache__/__init__.cpython-311.pyc +0 -0
  47. package/src/trace_engine/findings/__pycache__/correlate.cpython-311.pyc +0 -0
  48. package/src/trace_engine/findings/__pycache__/model.cpython-311.pyc +0 -0
  49. package/src/trace_engine/findings/__pycache__/recommendations.cpython-311.pyc +0 -0
  50. package/src/trace_engine/findings/__pycache__/store.cpython-311.pyc +0 -0
  51. package/src/trace_engine/findings/correlate.py +103 -0
  52. package/src/trace_engine/findings/model.py +40 -0
  53. package/src/trace_engine/findings/recommendations.py +35 -0
  54. package/src/trace_engine/findings/store.py +39 -0
  55. package/src/trace_engine/framework/__init__.py +58 -0
  56. package/src/trace_engine/framework/__pycache__/__init__.cpython-311.pyc +0 -0
  57. package/src/trace_engine/framework/__pycache__/base.cpython-311.pyc +0 -0
  58. package/src/trace_engine/framework/__pycache__/csharp.cpython-311.pyc +0 -0
  59. package/src/trace_engine/framework/__pycache__/dart.cpython-311.pyc +0 -0
  60. package/src/trace_engine/framework/__pycache__/django.cpython-311.pyc +0 -0
  61. package/src/trace_engine/framework/__pycache__/express.cpython-311.pyc +0 -0
  62. package/src/trace_engine/framework/__pycache__/fastapi.cpython-311.pyc +0 -0
  63. package/src/trace_engine/framework/__pycache__/flask.cpython-311.pyc +0 -0
  64. package/src/trace_engine/framework/__pycache__/go.cpython-311.pyc +0 -0
  65. package/src/trace_engine/framework/__pycache__/nextjs.cpython-311.pyc +0 -0
  66. package/src/trace_engine/framework/__pycache__/php.cpython-311.pyc +0 -0
  67. package/src/trace_engine/framework/__pycache__/react_router.cpython-311.pyc +0 -0
  68. package/src/trace_engine/framework/__pycache__/ruby.cpython-311.pyc +0 -0
  69. package/src/trace_engine/framework/__pycache__/rust.cpython-311.pyc +0 -0
  70. package/src/trace_engine/framework/__pycache__/springboot.cpython-311.pyc +0 -0
  71. package/src/trace_engine/framework/base.py +49 -0
  72. package/src/trace_engine/framework/csharp.py +111 -0
  73. package/src/trace_engine/framework/dart.py +152 -0
  74. package/src/trace_engine/framework/django.py +188 -0
  75. package/src/trace_engine/framework/express.py +82 -0
  76. package/src/trace_engine/framework/fastapi.py +135 -0
  77. package/src/trace_engine/framework/flask.py +108 -0
  78. package/src/trace_engine/framework/go.py +94 -0
  79. package/src/trace_engine/framework/nextjs.py +200 -0
  80. package/src/trace_engine/framework/php.py +98 -0
  81. package/src/trace_engine/framework/react_router.py +331 -0
  82. package/src/trace_engine/framework/ruby.py +69 -0
  83. package/src/trace_engine/framework/rust.py +101 -0
  84. package/src/trace_engine/framework/springboot.py +145 -0
  85. package/src/trace_engine/harness/__init__.py +19 -0
  86. package/src/trace_engine/harness/__pycache__/__init__.cpython-311.pyc +0 -0
  87. package/src/trace_engine/harness/__pycache__/benchmark.cpython-311.pyc +0 -0
  88. package/src/trace_engine/harness/__pycache__/context.cpython-311.pyc +0 -0
  89. package/src/trace_engine/harness/__pycache__/engine.cpython-311.pyc +0 -0
  90. package/src/trace_engine/harness/__pycache__/patcher.cpython-311.pyc +0 -0
  91. package/src/trace_engine/harness/__pycache__/remediators.cpython-311.pyc +0 -0
  92. package/src/trace_engine/harness/benchmark.py +55 -0
  93. package/src/trace_engine/harness/context.py +49 -0
  94. package/src/trace_engine/harness/engine.py +187 -0
  95. package/src/trace_engine/harness/patcher.py +143 -0
  96. package/src/trace_engine/harness/remediators.py +307 -0
  97. package/src/trace_engine/ingest/__init__.py +15 -0
  98. package/src/trace_engine/ingest/__pycache__/__init__.cpython-311.pyc +0 -0
  99. package/src/trace_engine/ingest/__pycache__/files.cpython-311.pyc +0 -0
  100. package/src/trace_engine/ingest/__pycache__/hashing.cpython-311.pyc +0 -0
  101. package/src/trace_engine/ingest/__pycache__/ignore.cpython-311.pyc +0 -0
  102. package/src/trace_engine/ingest/__pycache__/repository.cpython-311.pyc +0 -0
  103. package/src/trace_engine/ingest/files.py +60 -0
  104. package/src/trace_engine/ingest/hashing.py +21 -0
  105. package/src/trace_engine/ingest/ignore.py +108 -0
  106. package/src/trace_engine/ingest/repository.py +50 -0
  107. package/src/trace_engine/intelligence/__init__.py +15 -0
  108. package/src/trace_engine/intelligence/__pycache__/__init__.cpython-311.pyc +0 -0
  109. package/src/trace_engine/intelligence/__pycache__/evaluation.cpython-311.pyc +0 -0
  110. package/src/trace_engine/intelligence/__pycache__/orchestrator.cpython-311.pyc +0 -0
  111. package/src/trace_engine/intelligence/__pycache__/tracebench.cpython-311.pyc +0 -0
  112. package/src/trace_engine/intelligence/evaluation.py +828 -0
  113. package/src/trace_engine/intelligence/laya/__init__.py +19 -0
  114. package/src/trace_engine/intelligence/laya/__pycache__/__init__.cpython-311.pyc +0 -0
  115. package/src/trace_engine/intelligence/laya/__pycache__/prompts.cpython-311.pyc +0 -0
  116. package/src/trace_engine/intelligence/laya/__pycache__/router.cpython-311.pyc +0 -0
  117. package/src/trace_engine/intelligence/laya/__pycache__/schemas.cpython-311.pyc +0 -0
  118. package/src/trace_engine/intelligence/laya/__pycache__/telemetry.cpython-311.pyc +0 -0
  119. package/src/trace_engine/intelligence/laya/__pycache__/thresholds.cpython-311.pyc +0 -0
  120. package/src/trace_engine/intelligence/laya/prompts.py +67 -0
  121. package/src/trace_engine/intelligence/laya/router.py +352 -0
  122. package/src/trace_engine/intelligence/laya/schemas.py +56 -0
  123. package/src/trace_engine/intelligence/laya/telemetry.py +48 -0
  124. package/src/trace_engine/intelligence/laya/thresholds.py +12 -0
  125. package/src/trace_engine/intelligence/orchestrator.py +130 -0
  126. package/src/trace_engine/intelligence/securebert/__init__.py +6 -0
  127. package/src/trace_engine/intelligence/securebert/__pycache__/__init__.cpython-311.pyc +0 -0
  128. package/src/trace_engine/intelligence/securebert/__pycache__/cache.cpython-311.pyc +0 -0
  129. package/src/trace_engine/intelligence/securebert/__pycache__/classifier.cpython-311.pyc +0 -0
  130. package/src/trace_engine/intelligence/securebert/cache.py +37 -0
  131. package/src/trace_engine/intelligence/securebert/classifier.py +240 -0
  132. package/src/trace_engine/intelligence/tracebench.py +61 -0
  133. package/src/trace_engine/intelligence/training/__init__.py +21 -0
  134. package/src/trace_engine/intelligence/training/__pycache__/__init__.cpython-311.pyc +0 -0
  135. package/src/trace_engine/intelligence/training/__pycache__/dataset.cpython-311.pyc +0 -0
  136. package/src/trace_engine/intelligence/training/__pycache__/dataset_importers.cpython-311.pyc +0 -0
  137. package/src/trace_engine/intelligence/training/__pycache__/laya_trainer.cpython-311.pyc +0 -0
  138. package/src/trace_engine/intelligence/training/__pycache__/lora_system2.cpython-311.pyc +0 -0
  139. package/src/trace_engine/intelligence/training/__pycache__/morefixes_pipeline.cpython-311.pyc +0 -0
  140. package/src/trace_engine/intelligence/training/__pycache__/slicer.cpython-311.pyc +0 -0
  141. package/src/trace_engine/intelligence/training/__pycache__/train_all.cpython-311.pyc +0 -0
  142. package/src/trace_engine/intelligence/training/__pycache__/trainer.cpython-311.pyc +0 -0
  143. package/src/trace_engine/intelligence/training/dataset.py +320 -0
  144. package/src/trace_engine/intelligence/training/dataset_importers.py +349 -0
  145. package/src/trace_engine/intelligence/training/laya_trainer.py +678 -0
  146. package/src/trace_engine/intelligence/training/lora_system2.py +162 -0
  147. package/src/trace_engine/intelligence/training/morefixes_pipeline.py +463 -0
  148. package/src/trace_engine/intelligence/training/slicer.py +129 -0
  149. package/src/trace_engine/intelligence/training/train_all.py +1009 -0
  150. package/src/trace_engine/intelligence/training/trainer.py +321 -0
  151. package/src/trace_engine/interactive.py +623 -0
  152. package/src/trace_engine/mcp/__init__.py +6 -0
  153. package/src/trace_engine/mcp/__pycache__/__init__.cpython-311.pyc +0 -0
  154. package/src/trace_engine/mcp/__pycache__/config.cpython-311.pyc +0 -0
  155. package/src/trace_engine/mcp/__pycache__/server.cpython-311.pyc +0 -0
  156. package/src/trace_engine/mcp/config.py +42 -0
  157. package/src/trace_engine/mcp/server.py +398 -0
  158. package/src/trace_engine/output/__init__.py +28 -0
  159. package/src/trace_engine/output/__pycache__/__init__.cpython-311.pyc +0 -0
  160. package/src/trace_engine/output/__pycache__/html.cpython-311.pyc +0 -0
  161. package/src/trace_engine/output/__pycache__/markdown.cpython-311.pyc +0 -0
  162. package/src/trace_engine/output/__pycache__/sarif.cpython-311.pyc +0 -0
  163. package/src/trace_engine/output/__pycache__/terminal.cpython-311.pyc +0 -0
  164. package/src/trace_engine/output/html.py +142 -0
  165. package/src/trace_engine/output/markdown.py +65 -0
  166. package/src/trace_engine/output/sarif.py +180 -0
  167. package/src/trace_engine/output/terminal.py +537 -0
  168. package/src/trace_engine/parsing/__init__.py +26 -0
  169. package/src/trace_engine/parsing/__pycache__/__init__.cpython-311.pyc +0 -0
  170. package/src/trace_engine/parsing/__pycache__/calls.cpython-311.pyc +0 -0
  171. package/src/trace_engine/parsing/__pycache__/language.cpython-311.pyc +0 -0
  172. package/src/trace_engine/parsing/__pycache__/locations.cpython-311.pyc +0 -0
  173. package/src/trace_engine/parsing/__pycache__/parser.cpython-311.pyc +0 -0
  174. package/src/trace_engine/parsing/__pycache__/symbols.cpython-311.pyc +0 -0
  175. package/src/trace_engine/parsing/calls.py +14 -0
  176. package/src/trace_engine/parsing/language.py +24 -0
  177. package/src/trace_engine/parsing/locations.py +17 -0
  178. package/src/trace_engine/parsing/parser.py +465 -0
  179. package/src/trace_engine/parsing/symbols.py +45 -0
  180. package/src/trace_engine/plugin/__init__.py +157 -0
  181. package/src/trace_engine/plugin/__pycache__/__init__.cpython-311.pyc +0 -0
  182. package/src/trace_engine/plugin/__pycache__/evaluator.cpython-311.pyc +0 -0
  183. package/src/trace_engine/plugin/__pycache__/swebench_adapter.cpython-311.pyc +0 -0
  184. package/src/trace_engine/plugin/__pycache__/task.cpython-311.pyc +0 -0
  185. package/src/trace_engine/plugin/bundle/hooks.json +24 -0
  186. package/src/trace_engine/plugin/bundle/mcp_config.json +11 -0
  187. package/src/trace_engine/plugin/bundle/plugin.json +20 -0
  188. package/src/trace_engine/plugin/bundle/rules/security_remediation.md +40 -0
  189. package/src/trace_engine/plugin/bundle/skills/trace-security-harness/SKILL.md +118 -0
  190. package/src/trace_engine/plugin/evaluator.py +105 -0
  191. package/src/trace_engine/plugin/swebench_adapter.py +96 -0
  192. package/src/trace_engine/plugin/task.py +48 -0
  193. package/src/trace_engine/policy/__init__.py +5 -0
  194. package/src/trace_engine/policy/__pycache__/__init__.cpython-311.pyc +0 -0
  195. package/src/trace_engine/policy/__pycache__/scope.cpython-311.pyc +0 -0
  196. package/src/trace_engine/policy/scope.py +80 -0
  197. package/src/trace_engine/runtime/__init__.py +7 -0
  198. package/src/trace_engine/runtime/__pycache__/__init__.cpython-311.pyc +0 -0
  199. package/src/trace_engine/runtime/__pycache__/client.cpython-311.pyc +0 -0
  200. package/src/trace_engine/runtime/__pycache__/observations.cpython-311.pyc +0 -0
  201. package/src/trace_engine/runtime/__pycache__/target.cpython-311.pyc +0 -0
  202. package/src/trace_engine/runtime/client.py +81 -0
  203. package/src/trace_engine/runtime/observations.py +18 -0
  204. package/src/trace_engine/runtime/target.py +23 -0
  205. package/src/trace_engine/security/__init__.py +16 -0
  206. package/src/trace_engine/security/__pycache__/__init__.cpython-311.pyc +0 -0
  207. package/src/trace_engine/security/__pycache__/fusion.cpython-311.pyc +0 -0
  208. package/src/trace_engine/security/__pycache__/hypotheses.cpython-311.pyc +0 -0
  209. package/src/trace_engine/security/__pycache__/signals.cpython-311.pyc +0 -0
  210. package/src/trace_engine/security/__pycache__/timing.cpython-311.pyc +0 -0
  211. package/src/trace_engine/security/fusion.py +126 -0
  212. package/src/trace_engine/security/hypotheses.py +251 -0
  213. package/src/trace_engine/security/signals.py +28 -0
  214. package/src/trace_engine/security/timing.py +132 -0
  215. package/src/trace_engine/testpacks/__init__.py +18 -0
  216. package/src/trace_engine/testpacks/__pycache__/__init__.cpython-311.pyc +0 -0
  217. package/src/trace_engine/testpacks/__pycache__/authentication.cpython-311.pyc +0 -0
  218. package/src/trace_engine/testpacks/__pycache__/base.cpython-311.pyc +0 -0
  219. package/src/trace_engine/testpacks/__pycache__/bfla.cpython-311.pyc +0 -0
  220. package/src/trace_engine/testpacks/__pycache__/bola.cpython-311.pyc +0 -0
  221. package/src/trace_engine/testpacks/__pycache__/cors.cpython-311.pyc +0 -0
  222. package/src/trace_engine/testpacks/__pycache__/deserialization.cpython-311.pyc +0 -0
  223. package/src/trace_engine/testpacks/__pycache__/injection.cpython-311.pyc +0 -0
  224. package/src/trace_engine/testpacks/__pycache__/mass_assignment.cpython-311.pyc +0 -0
  225. package/src/trace_engine/testpacks/__pycache__/path_traversal.cpython-311.pyc +0 -0
  226. package/src/trace_engine/testpacks/__pycache__/registry.cpython-311.pyc +0 -0
  227. package/src/trace_engine/testpacks/__pycache__/ssrf.cpython-311.pyc +0 -0
  228. package/src/trace_engine/testpacks/__pycache__/ssti.cpython-311.pyc +0 -0
  229. package/src/trace_engine/testpacks/authentication.py +60 -0
  230. package/src/trace_engine/testpacks/base.py +63 -0
  231. package/src/trace_engine/testpacks/bfla.py +70 -0
  232. package/src/trace_engine/testpacks/bola.py +94 -0
  233. package/src/trace_engine/testpacks/cors.py +85 -0
  234. package/src/trace_engine/testpacks/deserialization.py +86 -0
  235. package/src/trace_engine/testpacks/injection.py +179 -0
  236. package/src/trace_engine/testpacks/mass_assignment.py +70 -0
  237. package/src/trace_engine/testpacks/path_traversal.py +117 -0
  238. package/src/trace_engine/testpacks/registry.py +44 -0
  239. package/src/trace_engine/testpacks/ssrf.py +85 -0
  240. package/src/trace_engine/testpacks/ssti.py +96 -0
  241. package/src/trace_engine/tools/__init__.py +6 -0
  242. package/src/trace_engine/tools/__pycache__/__init__.cpython-311.pyc +0 -0
  243. package/src/trace_engine/tools/__pycache__/adapters.cpython-311.pyc +0 -0
  244. package/src/trace_engine/tools/__pycache__/base.cpython-311.pyc +0 -0
  245. package/src/trace_engine/tools/__pycache__/registry.cpython-311.pyc +0 -0
  246. package/src/trace_engine/tools/adapters.py +111 -0
  247. package/src/trace_engine/tools/base.py +68 -0
  248. package/src/trace_engine/tools/registry.py +36 -0
  249. 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