competition-vd/scripts/run_training.py

102 lines
4.5 KiB
Python
Executable File

import sys
import os
sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
from src.models.linevul import LineVulModel, LineVulTrainer
from src.models.xgboost_classifier import XGBoostVulClassifier, XGBoostVulRegressor
from src.data.data_loader import DataLoader
from src.evaluation.evaluator import Evaluator
import argparse
import torch
import numpy as np
from sklearn.model_selection import train_test_split
def train_linevul(args, train_data, val_data):
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
model = LineVulModel(num_labels=len(set(train_data['labels'])))
trainer = LineVulTrainer(model, device, learning_rate=args.learning_rate)
train_dataset = torch.utils.data.TensorDataset(
torch.tensor(train_data['features']).float(),
torch.tensor(train_data['labels']).long()
)
val_dataset = torch.utils.data.TensorDataset(
torch.tensor(val_data['features']).float(),
torch.tensor(val_data['labels']).long()
)
train_loader = torch.utils.data.DataLoader(train_dataset, batch_size=args.batch_size, shuffle=True)
val_loader = torch.utils.data.DataLoader(val_dataset, batch_size=args.batch_size)
trainer.train(train_loader, args.num_epochs)
_, preds, labels = trainer.evaluate(val_loader)
return model, preds, labels
def train_xgboost(args, train_data, val_data):
classifier = XGBoostVulClassifier(
n_estimators=args.n_estimators,
max_depth=args.max_depth,
learning_rate=args.learning_rate
)
classifier.train(train_data['features'], train_data['labels'])
regressor = XGBoostVulRegressor(
n_estimators=args.n_estimators,
max_depth=args.max_depth,
learning_rate=args.learning_rate
)
regressor.train(train_data['features'], train_data['cvss_scores'])
class_preds = classifier.predict(val_data['features'])
reg_preds = regressor.predict(val_data['features'])
return classifier, regressor, class_preds, reg_preds
def main(args):
data_loader = DataLoader(args.data_dir)
train_data = {
'features': np.load(os.path.join(args.data_dir, 'train', 'features.npy')),
'labels': np.load(os.path.join(args.data_dir, 'train', 'labels.npy')),
'cvss_scores': np.load(os.path.join(args.data_dir, 'train', 'cvss_scores.npy'))
}
val_data = {
'features': np.load(os.path.join(args.data_dir, 'valid', 'features.npy')),
'labels': np.load(os.path.join(args.data_dir, 'valid', 'labels.npy')),
'cvss_scores': np.load(os.path.join(args.data_dir, 'valid', 'cvss_scores.npy'))
}
if args.model == 'linevul':
model, preds, labels = train_linevul(args, train_data, val_data)
torch.save(model.state_dict(), os.path.join(args.output_dir, 'linevul_model.pth'))
elif args.model == 'xgboost':
classifier, regressor, class_preds, reg_preds = train_xgboost(args, train_data, val_data)
classifier.model.save_model(os.path.join(args.output_dir, 'xgboost_classifier.json'))
regressor.model.save_model(os.path.join(args.output_dir, 'xgboost_regressor.json'))
else:
raise ValueError(f"Unsupported model: {args.model}")
evaluator = Evaluator()
if args.model == 'linevul':
results = evaluator.evaluate_classification(labels, preds)
else:
results = evaluator.evaluate_classification(val_data['labels'], class_preds)
results.update(evaluator.evaluate_regression(val_data['cvss_scores'], reg_preds))
evaluator.evaluate_and_print(results)
if __name__ == "__main__":
parser = argparse.ArgumentParser(description="Run training for vulnerability detection models")
parser.add_argument('--model', type=str, choices=['linevul', 'xgboost'], required=True, help='Model to train')
parser.add_argument('--data_dir', type=str, default='data/processed', help='Directory containing processed data')
parser.add_argument('--output_dir', type=str, default='outputs/models', help='Directory to save trained models')
parser.add_argument('--batch_size', type=int, default=32, help='Batch size for training')
parser.add_argument('--num_epochs', type=int, default=10, help='Number of epochs for training')
parser.add_argument('--learning_rate', type=float, default=2e-5, help='Learning rate')
parser.add_argument('--n_estimators', type=int, default=100, help='Number of estimators for XGBoost')
parser.add_argument('--max_depth', type=int, default=3, help='Max depth for XGBoost')
args = parser.parse_args()
main(args)