File size: 4,039 Bytes
fe66586
 
 
 
 
 
 
 
 
 
2d9d6c4
 
2607de5
 
 
 
 
 
 
 
 
fe66586
2d9d6c4
fe66586
 
2607de5
2d9d6c4
fe66586
 
2607de5
fe66586
 
 
 
 
 
2607de5
2d9d6c4
 
 
 
 
 
 
 
 
2607de5
2d9d6c4
 
 
2607de5
2d9d6c4
 
 
 
 
 
fe66586
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
import matplotlib.pyplot as plt
import seaborn as sns
import numpy as np
import os
from sklearn.metrics import confusion_matrix, roc_curve, auc, classification_report
from sklearn.preprocessing import label_binarize
from . import config


def plot_confusion_matrix(y_true, y_pred, classes, model_name):
    """Generates and saves both Raw and Normalized confusion matrix heatmaps."""

    # 1. Determine Alphabetical Order
    # y_true/y_pred are indices (0, 1, 2...) mapping to the original 'classes' list.
    # We want to display them in alphabetical order of the class names.
    sorted_indices = np.argsort(classes)
    sorted_classes = np.array(classes)[sorted_indices]

    # 2. Calculate Raw Matrix
    # passing 'labels=sorted_indices' forces the matrix rows/cols to follow the alphabetical order
    cm = confusion_matrix(y_true, y_pred, labels=sorted_indices)

    # Plot Raw Counts
    plt.figure(figsize=(10, 8))
    sns.heatmap(cm, annot=True, fmt='d', cmap='Blues',
                xticklabels=sorted_classes, yticklabels=sorted_classes)
    plt.title(f"{model_name} Confusion Matrix (Counts)")
    plt.ylabel('True Label')
    plt.xlabel('Predicted Label')
    plt.xticks(rotation=45)
    plt.tight_layout()

    filename = f"{model_name.lower()}_confusion_matrix.png"
    plt.savefig(os.path.join(config.MODEL_DIR, filename))
    plt.close()

    # 3. Calculate Normalized Matrix
    # Divide each row element by the sum of that row (True Label count)
    cm_norm = cm.astype('float') / cm.sum(axis=1)[:, np.newaxis]
    # Replace NaN with 0 (safe guard for empty classes)
    cm_norm = np.nan_to_num(cm_norm)

    # Plot Normalized Percentages
    plt.figure(figsize=(10, 8))
    # Use fmt='.2f' to show 2 decimal places (e.g., 0.95)
    sns.heatmap(cm_norm, annot=True, fmt='.2f', cmap='Greens',
                xticklabels=sorted_classes, yticklabels=sorted_classes)
    plt.title(f"{model_name} Confusion Matrix (Normalized)")
    plt.ylabel('True Label')
    plt.xlabel('Predicted Label')
    plt.xticks(rotation=45)
    plt.tight_layout()

    filename_norm = f"{model_name.lower()}_confusion_matrix_normalized.png"
    plt.savefig(os.path.join(config.MODEL_DIR, filename_norm))
    plt.close()


def plot_multiclass_roc(model, X_test, y_test, classes, model_name):
    """Generates and saves a Multi-class ROC Curve."""
    # 1. Binarize labels (One-vs-Rest)
    y_test_bin = label_binarize(y_test, classes=range(len(classes)))
    n_classes = y_test_bin.shape[1]

    # 2. Get probabilities
    if hasattr(model, "predict_proba"):
        y_score = model.predict_proba(X_test)
    else:
        print(f"{model_name} does not support probability prediction. Skipping ROC.")
        return

    # 3. Compute ROC curve and ROC area for each class
    fpr = dict()
    tpr = dict()
    roc_auc = dict()

    for i in range(n_classes):
        fpr[i], tpr[i], _ = roc_curve(y_test_bin[:, i], y_score[:, i])
        roc_auc[i] = auc(fpr[i], tpr[i])

    # 4. Plot
    plt.figure(figsize=(10, 8))
    colors = plt.cm.rainbow(np.linspace(0, 1, n_classes))

    for i, color in zip(range(n_classes), colors):
        plt.plot(fpr[i], tpr[i], color=color, lw=2,
                 label=f'{classes[i]} (AUC = {roc_auc[i]:.2f})')

    plt.plot([0, 1], [0, 1], 'k--', lw=2)  # Diagonal line
    plt.xlim([0.0, 1.0])
    plt.ylim([0.0, 1.05])
    plt.xlabel('False Positive Rate')
    plt.ylabel('True Positive Rate')
    plt.title(f'{model_name} Multi-class ROC Curve')
    plt.legend(loc="lower right")
    plt.tight_layout()

    filename = f"{model_name.lower()}_roc_curve.png"
    plt.savefig(os.path.join(config.MODEL_DIR, filename))
    plt.close()


def save_classification_report(y_true, y_pred, classes, model_name):
    """Saves the text classification report."""
    report = classification_report(y_true, y_pred, target_names=classes)
    filename = f"{model_name.lower()}_report.txt"
    with open(os.path.join(config.MODEL_DIR, filename), "w") as f:
        f.write(report)
    print(f"Report saved to {filename}")