clasificador de severidad de vulnerabilidades (cve) usando distilbert + pytorch.
dado el texto de descripción de un cve, el modelo predice su nivel de severidad: LOW, MEDIUM, HIGH o CRITICAL.
- uv instalado
- python 3.12+
- acceso a internet para descargar datos y el modelo base de huggingface
git clone https://github.com/tu-usuario/cve-classifier
cd cve-classifier
uv venv
source .venv/bin/activate
uv synccrea un archivo .env en la raíz del proyecto:
NVD_API_KEY=tu-api-key-aquipuedes obtener una api key gratuita en: https://nvd.nist.gov/developers/request-an-api-key
sin api key el script funciona igual pero con rate limit más estricto (0.6s entre requests).
# -- 1. descargar datos
make fetch
# -- 2. preprocesar datos
make preprocess
# -- 3. entrenar modelo
make train
# -- 4. clasificar cves nuevos
make inferencedescarga cves desde la api pública de nist nvd. descarga 10,000 cves generales + 2,000 cves de severidad critical para balancear clases.
make fetch
# output: data/cves_raw.jsondistribución típica después del fetch:
HIGH: ~4500
MEDIUM: ~4400
CRITICAL: ~2000
LOW: ~800
limpia las descripciones, convierte severidad a label numérico y genera el split train/val/test (80/10/10).
make preprocess
# output: data/processed/train.csv
# data/processed/val.csv
# data/processed/test.csv
# data/processed/label_map.jsonlabel map:
{
"LOW": 0,
"MEDIUM": 1,
"HIGH": 2,
"CRITICAL": 3
}fine-tune de distilbert-base-uncased para clasificación de 4 clases.
usa class weights para manejar el desbalance entre clases.
detecta automáticamente si hay gpu (mps en apple silicon, cuda en nvidia).
make train
# output: models/checkpoints/epoch_1.pt
# models/checkpoints/epoch_2.pt
# models/checkpoints/epoch_3.pt
# models/checkpoints/best_model.ptresultados después de 3 epochs:
| clase | precision | recall | f1 |
|---|---|---|---|
| LOW | 0.40 | 0.81 | 0.53 |
| MEDIUM | 0.79 | 0.66 | 0.72 |
| HIGH | 0.78 | 0.78 | 0.78 |
| CRITICAL | 0.95 | 0.89 | 0.91 |
| accuracy | 0.75 |
clasifica descripciones de cves nuevos usando el modelo entrenado.
make inferenceejemplo de output:
description: buffer overflow in the kernel allows local attackers to execute arbitrary code...
severity: HIGH (confidence: 0.9628)
probs: {'LOW': 0.005, 'MEDIUM': 0.0291, 'HIGH': 0.9628, 'CRITICAL': 0.0031}
description: remote code execution vulnerability in apache allows unauthenticated attackers...
severity: CRITICAL (confidence: 0.6038)
probs: {'LOW': 0.0025, 'MEDIUM': 0.0253, 'HIGH': 0.3683, 'CRITICAL': 0.6038}
también puedes usar el predictor directamente en tu código:
from inference.predict import load_model, predict
from models.classifier import get_tokenizer
tokenizer = get_tokenizer()
model = load_model()
result = predict("sql injection allows remote attacker to dump the database", model, tokenizer)
print(result)
# {'severity': 'HIGH', 'confidence': 0.87, 'probabilities': {...}}cve-classifier/
├── data/
│ ├── __init__.py
│ ├── fetch_cves.py # descarga cves desde nist nvd api
│ ├── preprocess.py # limpieza y split de datos
│ └── processed/
│ ├── train.csv
│ ├── val.csv
│ ├── test.csv
│ └── label_map.json
├── models/
│ ├── __init__.py
│ ├── classifier.py # arquitectura distilbert + clasificador
│ ├── dataset.py # pytorch dataset
│ └── checkpoints/
│ └── best_model.pt
├── training/
│ ├── __init__.py
│ └── train.py # loop de entrenamiento
├── inference/
│ ├── __init__.py
│ └── predict.py # inferencia sobre cves nuevos
├── .env # variables de entorno (no commitear)
├── .gitignore
├── Makefile
├── pyproject.toml
└── README.md
| paquete | uso |
|---|---|
| torch | framework de deep learning |
| transformers | modelo distilbert de huggingface |
| pandas | manejo de datos |
| scikit-learn | split y métricas |
| requests | fetch a la api de nist |
| python-dotenv | manejo de variables de entorno |
- el modelo base
distilbert-base-uncasedse descarga automáticamente desde huggingface la primera vez - los checkpoints se guardan en
models/checkpoints/después de cada epoch - la clase
LOWtiene el f1 más bajo (0.53) por menor cantidad de datos; mejorable descargando más cves low con la api key activa - para reentrenar desde cero borra
models/checkpoints/y corremake train
.venv/
.env
data/cves_raw.json
data/processed/
models/checkpoints/
__pycache__/
*.pyc