forked from Eshe/competition-vd
102 lines
4.5 KiB
Python
Executable File
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) |