File size: 3,962 Bytes
168ae1c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
import torch
from transformers import (
    AutoTokenizer, 
    AutoModelForSequenceClassification,
    Trainer, 
    TrainingArguments
)
from datasets import load_from_disk
import numpy as np
from sklearn.metrics import accuracy_score, precision_recall_fscore_support
import os

print("="*60)
print("RETRAINING WITH ENHANCED DATASET")
print("="*60)

# Load enhanced dataset
print("πŸ“‚ Loading enhanced dataset...")
dataset = load_from_disk("enhanced_data/hf_dataset")

print(f"Train size: {len(dataset['train'])}")
print(f"Test size: {len(dataset['test'])}")

# Load tokenizer and model
print("πŸ€– Loading CodeBERT...")
model_name = "microsoft/codebert-base"
tokenizer = AutoTokenizer.from_pretrained(model_name)
model = AutoModelForSequenceClassification.from_pretrained(
    model_name, 
    num_labels=2
)

# Tokenization function
def tokenize_function(examples):
    return tokenizer(
        examples["code"],
        padding="max_length",
        truncation=True,
        max_length=256  # Increased for better context
    )

print("πŸ”’ Tokenizing dataset...")
tokenized_datasets = dataset.map(tokenize_function, batched=True)

# Remove text columns
columns_to_remove = ["code", "type", "explanation", "has_syntax_error", "syntax_error"]
columns_to_remove = [col for col in columns_to_remove if col in tokenized_datasets["train"].column_names]
tokenized_datasets = tokenized_datasets.remove_columns(columns_to_remove)
tokenized_datasets.set_format("torch")

# Training arguments (better this time)
training_args = TrainingArguments(
    output_dir="./enhanced_model",
    num_train_epochs=5,  # More epochs
    per_device_train_batch_size=16,
    per_device_eval_batch_size=16,
    warmup_steps=500,
    weight_decay=0.01,
    logging_dir="./enhanced_logs",
    logging_steps=50,
    eval_strategy="epoch",
    save_strategy="epoch",
    load_best_model_at_end=True,
    metric_for_best_model="f1",
    greater_is_better=True,
    save_total_limit=2,
    report_to="none"
)

# Better metrics
def compute_metrics(p):
    predictions, labels = p
    predictions = np.argmax(predictions, axis=1)
    
    accuracy = accuracy_score(labels, predictions)
    precision, recall, f1, _ = precision_recall_fscore_support(
        labels, predictions, average="binary", zero_division=0
    )
    
    return {
        "accuracy": accuracy,
        "precision": precision,
        "recall": recall,
        "f1": f1
    }

# Create trainer
trainer = Trainer(
    model=model,
    args=training_args,
    train_dataset=tokenized_datasets["train"],
    eval_dataset=tokenized_datasets["test"],
    compute_metrics=compute_metrics,
)

# Train
print("πŸš€ Training enhanced model...")
print("This will take 10-15 minutes...")
trainer.train()

# Evaluate
print("\nπŸ“Š Final Evaluation:")
metrics = trainer.evaluate()
print(f"Accuracy: {metrics['eval_accuracy']:.2%}")
print(f"Precision: {metrics['eval_precision']:.2%}")
print(f"Recall: {metrics['eval_recall']:.2%}")
print(f"F1 Score: {metrics['eval_f1']:.2%}")

# Save model
print("\nπŸ’Ύ Saving enhanced model...")
trainer.save_model("enhanced_saved_model")
tokenizer.save_pretrained("enhanced_saved_model")

print("\n" + "="*60)
print("πŸŽ‰ ENHANCED MODEL TRAINED!")
print("Model saved to: enhanced_saved_model/")
print("="*60)

# Quick test
print("\nπŸ” Quick test:")
test_codes = [
    """query = f"SELECT * FROM users WHERE id = {user_id}" """,
    """api_key = os.getenv("API_KEY")""",
    """def test()\n    print("hello")""",  # Syntax error
]

for code in test_codes:
    inputs = tokenizer(code, return_tensors="pt", truncation=True, max_length=256)
    with torch.no_grad():
        outputs = model(**inputs)
        probs = torch.nn.functional.softmax(outputs.logits, dim=-1)
    
    prediction = "VULNERABLE" if probs[0][1] > 0.5 else "SAFE"
    print(f"\nCode: {code[:50]}...")
    print(f"  Prediction: {prediction}")
    print(f"  Confidence: Safe={probs[0][0]:.2%}, Vulnerable={probs[0][1]:.2%}")