Spaces:
Running
Running
| from datasets import load_dataset | |
| import pandas as pd | |
| import matplotlib.pyplot as plt | |
| from sklearn.model_selection import train_test_split | |
| from sklearn.preprocessing import StandardScaler | |
| from sklearn.neighbors import KNeighborsClassifier | |
| from sklearn.metrics import accuracy_score | |
| print("1. Downloading dataset from Hugging Face...") | |
| # Fetching a clean tabular dataset directly from Hugging Face Hub | |
| dataset = load_dataset("scikit-learn/iris", split="train") | |
| df = pd.DataFrame(dataset) | |
| # Features (Measurements) and Target (Species) | |
| X = df[['sepal length (cm)', 'sepal width (cm)', 'petal length (cm)', 'petal width (cm)']] | |
| y = df['target'] | |
| print("2. Splitting and scaling data...") | |
| # Split into training and testing sets | |
| X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.2, random_state=42) | |
| # Scale features for accurate KNN distance calculations | |
| scaler = StandardScaler() | |
| X_train_scaled = scaler.fit_transform(X_train) | |
| X_test_scaled = scaler.transform(X_test) | |
| print("3. Training the KNN Model...") | |
| # Train KNN with 3 neighbors | |
| knn = KNeighborsClassifier(n_neighbors=3) | |
| knn.fit(X_train_scaled, y_train) | |
| # Check model accuracy | |
| y_pred = knn.predict(X_test_scaled) | |
| print(f"Model Accuracy: {accuracy_score(y_test, y_pred) * 100:.2f}%") | |
| # Test a brand new sample prediction | |
| new_sample = [[5.1, 3.5, 1.4, 0.2]] | |
| new_sample_scaled = scaler.transform(new_sample) | |
| prediction = knn.predict(new_sample_scaled) | |
| print(f"Prediction for new sample [Class]: {prediction[0]}") | |
| print("4. Generating Graph...") | |
| # Simple 2D plot using first two features (Sepal Length vs Sepal Width) | |
| plt.figure(figsize=(8, 6)) | |
| plt.scatter(df['sepal length (cm)'], df['sepal width (cm)'], c=df['target'], cmap='viridis', s=80, edgecolors='k') | |
| plt.title("Hugging Face Dataset - KNN Classification (Iris Sepal Dimensions)") | |
| plt.xlabel("Sepal Length (cm)") | |
| plt.ylabel("Sepal Width (cm)") | |
| plt.grid(True) | |
| plt.show() |