Team Ai
Datasetpublic

HighFive-OPJ/Implementation_1-Machine_learning

sourceHugging Faceupdated 1y agoView on Hugging Face
0likes21downloads
implementation1.py85 linesDownload Raw Back to root
1import pandas as pd
2import matplotlib.pyplot as plt
3import csv
4from sklearn.model_selection import train_test_split, GridSearchCV
5from sklearn.svm import SVC
6from sklearn.neighbors import KNeighborsClassifier
7from sklearn.feature_extraction.text import TfidfVectorizer
8from sklearn.metrics import precision_score, recall_score, f1_score, accuracy_score, confusion_matrix
9
10file_path = 'Test-3.tsv'
11data = pd.read_csv(file_path, sep="\t", names=["Sentence", "Label"], skiprows=1, quoting=csv.QUOTE_NONE, encoding="utf-8")
12data.columns = data.columns.str.strip()
13
14data = data.dropna(subset=['Sentence', 'Label'])
15data['Sentence'] = data['Sentence'].astype(str)
16
17X = data['Sentence']
18y = data['Label']
19
20vectorizer = TfidfVectorizer(ngram_range=(1, 2), max_features=5000)
21X_tfidf = vectorizer.fit_transform(X)
22
23X_train, X_test, y_train, y_test = train_test_split(
24    X_tfidf, y, test_size=0.3, random_state=42, stratify=y)
25
26plt.figure(figsize=(8, 6))
27y.value_counts().sort_index().plot(kind='bar', color='skyblue')
28plt.title('Class Distribution')
29plt.xlabel('Class')
30plt.ylabel('Frequency')
31plt.xticks(rotation=0)
32plt.tight_layout()
33plt.savefig('class_distribution.png')
34plt.close()
35
36svm_model = SVC(kernel='rbf', degree=3, random_state=42, class_weight='balanced')
37svm_model.fit(X_train, y_train)
38svm_predictions = svm_model.predict(X_test)
39
40print("SVM Model Performance:")
41print(f"Precision: {precision_score(y_test, svm_predictions, average='weighted', zero_division=0):.4f}")
42print(f"Recall:    {recall_score(y_test, svm_predictions, average='weighted', zero_division=0):.4f}")
43print(f"F1-Score:  {f1_score(y_test, svm_predictions, average='weighted', zero_division=0):.4f}")
44print(f"Accuracy:  {accuracy_score(y_test, svm_predictions):.4f}")
45
46param_grid = {'n_neighbors': list(range(3, 21, 2))}
47knn = KNeighborsClassifier()
48grid_search = GridSearchCV(knn, param_grid, cv=5, scoring='accuracy')
49grid_search.fit(X_train, y_train)
50
51best_knn = grid_search.best_estimator_
52knn_predictions = best_knn.predict(X_test)
53
54print("\nKNN Model Performance:")
55print(f"Best k:    {grid_search.best_params_['n_neighbors']}")
56print(f"Precision: {precision_score(y_test, knn_predictions, average='weighted', zero_division=0):.4f}")
57print(f"Recall:    {recall_score(y_test, knn_predictions, average='weighted', zero_division=0):.4f}")
58print(f"F1-Score:  {f1_score(y_test, knn_predictions, average='weighted', zero_division=0):.4f}")
59print(f"Accuracy:  {accuracy_score(y_test, knn_predictions):.4f}")
60
61def plot_conf_matrix(cm, title, filename):
62    plt.figure(figsize=(6, 5))
63    plt.imshow(cm, interpolation='nearest', cmap=plt.cm.Blues)
64    plt.title(title)
65    plt.colorbar()
66    tick_marks = range(len(cm))
67    plt.xticks(tick_marks, tick_marks)
68    plt.yticks(tick_marks, tick_marks)
69    plt.xlabel('Predicted Label')
70    plt.ylabel('True Label')
71
72    thresh = cm.max() / 2.
73    for i in range(cm.shape[0]):
74        for j in range(cm.shape[1]):
75            plt.text(j, i, str(cm[i, j]),
76                     ha='center', va='center',
77                     color='white' if cm[i, j] > thresh else 'black')
78
79    plt.tight_layout()
80    plt.savefig(filename)
81    plt.close()
82
83plot_conf_matrix(confusion_matrix(y_test, svm_predictions), 'SVM Confusion Matrix', 'svm_conf_matrix.png')
84plot_conf_matrix(confusion_matrix(y_test, knn_predictions), 'KNN Confusion Matrix', 'knn_conf_matrix.png')
85