Curso
Vas a implementar KNN en el famoso conjunto de datos de Iris.
Nota: Puedes considerar hacer el curso de Machine Learning con Python o, si quieres contexto sobre cómo ha evolucionado el ML y mucho más, leer esta publicación.
Introducción
El aprendizaje automático surge de la informática y estudia principalmente el diseño de algoritmos que pueden aprender de la experiencia. Para aprender, necesitan datos con ciertos atributos a partir de los cuales los algoritmos tratan de encontrar patrones predictivos con sentido. A grandes rasgos, las tareas de ML pueden clasificarse como aprendizaje de conceptos, clustering, modelado predictivo, etc. El objetivo final de los algoritmos de ML es tomar decisiones correctas sin intervención humana. Predecir la bolsa o el tiempo son un par de aplicaciones típicas de los algoritmos de aprendizaje automático.
Existen diversos algoritmos de ML como árboles de decisión, Naive Bayes, Random forest, máquinas de soporte vectorial, k vecinos más cercanos, k-means clustering, etc.
Del conjunto de algoritmos de aprendizaje automático, el que vas a usar hoy es k-vecinos más cercanos.
Ahora bien, ¿qué es exactamente el algoritmo de k vecinos más cercanos? ¡Vamos a verlo!
¿Qué es k-vecinos más cercanos?
KNN o k-vecinos más cercanos es un algoritmo de aprendizaje supervisado; supervisado significa que utiliza las etiquetas de clase de los datos de entrenamiento durante la fase de aprendizaje. Es un algoritmo basado en instancias: clasifica nuevos puntos de datos en función de instancias almacenadas y etiquetadas (puntos de datos). KNN puede usarse tanto para clasificación como para regresión; sin embargo, su uso es más extendido en clasificación.
La k de KNN es una variable crucial, también conocida como hiperparámetro, que ayuda a clasificar un punto de datos con precisión. Más concretamente, k es el número de vecinos más cercanos de los que quieres obtener un voto al clasificar un nuevo punto.

Como ves, al aumentar el valor de k de 1 a 7, la frontera de decisión entre dos clases con ciertos puntos de datos se vuelve más suave.
La gran pregunta es: ¿cómo ocurre esta "magia", de modo que cada vez que llega un nuevo punto de datos se clasifica en función de los puntos almacenados?
Vamos a entenderlo rápidamente de esta manera:
- Primero, cargas todos los datos e inicializas el valor de k,
- Luego, calculas la distancia entre los puntos almacenados y el nuevo punto que quieres clasificar, usando distintas métricas de similitud o distancia: distancia Manhattan (L1), Euclídea (L2), similitud del coseno, distancia de Bhattacharyya, distancia de Chebyshev, etc.
- Después, ordenas los valores de distancia en orden descendente o ascendente y determinas los k vecinos más cercanos (superiores o inferiores según el orden).
- Recoges las etiquetas de los k vecinos más cercanos y utilizas una votación mayoritaria o ponderada para clasificar el nuevo punto. Se asigna una etiqueta de clase en función del punto de datos que obtenga la puntuación más alta entre los almacenados.
- Por último, devuelves la clase predicha para la nueva instancia.
La predicción puede ser de dos tipos: clasificación, en la que asignas una etiqueta de clase al nuevo punto, o regresión, en la que asignas un valor. A diferencia de la clasificación, en regresión se asigna al nuevo punto la media de los k vecinos más cercanos.
Desventajas de KNN: en primer lugar, la complejidad de buscar los vecinos más cercanos para cada nuevo punto. En segundo lugar, determinar el valor óptimo de k puede ser tedioso. Y, por último, no siempre está claro qué métrica de distancia usar para calcular los vecinos.
Suficiente teoría, ¿no? Vamos a cargar, analizar y entender los datos que vas a usar en este pequeño tutorial.
Cargar los datos de Iris
El conjunto de datos Iris consta de 150 muestras con tres clases: Iris-Setosa, Iris-Versicolor e Iris-Virginica. Cuatro características/atributos permiten identificar cada muestra como una de las tres clases: sepal-length, sepal-width, petal-length y petal-width.
Si lo prefieres, usa otro conjunto de datos público o uno privado.
Sklearn es una biblioteca de Python para aprendizaje automático muy utilizada en tareas de ciencia de datos. Incluye diversos algoritmos de clasificación, regresión y clustering, como support vector machines, random forests, gradient boosting, k-means, KNN, etc. Dentro de sklearn tienes la librería datasets con múltiples conjuntos de datos listos para usar, incluido Iris. Es bastante intuitivo y directo. Vamos a cargar el conjunto iris.
from sklearn.datasets import load_iris
load_iris incluye tanto los datos como las etiquetas de clase de cada muestra. Vamos a extraerlo.
data = load_iris().data
La variable data será un array de NumPy con forma (150,4): 150 muestras, cada una con cuatro atributos. Cada clase tiene 50 muestras.
data.shape
(150, 4)
Extraigamos ahora las etiquetas de clase.
labels = load_iris().target
labels.shape
(150,)
A continuación, necesitas combinar los datos y las etiquetas. Para ello, usarás la excelente librería de Python NumPy. NumPy añade soporte para arrays y matrices grandes y multidimensionales, junto con una amplia colección de funciones matemáticas de alto nivel para operar sobre ellos. ¡Importémosla!
import numpy as np
Como data es un array 2D, tendrás que redimensionar labels también a 2D.
labels = np.reshape(labels,(150,1))
Ahora usarás la función concatenate de numpy con axis=-1 para concatenar por la segunda dimensión.
data = np.concatenate([data,labels],axis=-1)
data.shape
(150, 5)
Seguidamente, importa la librería de análisis de datos de Python, pandas, útil para organizar los datos en formato tabular y realizar operaciones y manipulaciones. En particular, ofrece estructuras de datos y operaciones para manejar tablas numéricas y series temporales.
En este tutorial, usarás pandas bastante.
import pandas as pd
names = ['sepal-length', 'sepal-width', 'petal-length', 'petal-width', 'species']
dataset = pd.DataFrame(data,columns=names)
Ya tienes el data frame dataset con los datos y las etiquetas de clase que necesitas.
Antes de seguir, recuerda que la variable labels tiene las etiquetas como valores numéricos, pero las vas a convertir a nombres de flores o especies.
Para ello, seleccionarás solo la columna class y sustituirás cada uno de los tres valores numéricos por la especie correspondiente. Usarás inplace=True para modificar el data frame dataset.
dataset['species'].replace(0, 'Iris-setosa',inplace=True)
dataset['species'].replace(1, 'Iris-versicolor',inplace=True)
dataset['species'].replace(2, 'Iris-virginica',inplace=True)
Imprime las cinco primeras filas de dataset para ver cómo queda.
dataset.head(5)
| sepal-length | sepal-width | petal-length | petal-width | species | |
|---|---|---|---|---|---|
| 0 | 5.1 | 3.5 | 1.4 | 0.2 | Iris-setosa |
| 1 | 4.9 | 3.0 | 1.4 | 0.2 | Iris-setosa |
| 2 | 4.7 | 3.2 | 1.3 | 0.2 | Iris-setosa |
| 3 | 4.6 | 3.1 | 1.5 | 0.2 | Iris-setosa |
| 4 | 5.0 | 3.6 | 1.4 | 0.2 | Iris-setosa |
Analiza tus datos
Veamos rápidamente cómo son las tres flores al visualizarlas y en qué se diferencian, no solo en números sino también en la realidad.
(Fuente)Visualicemos los datos que cargaste arriba con un scatterplot para ver cuánto afecta una variable a otra, o, dicho de otro modo, qué grado de correlación hay entre ambas.
Usarás la librería matplotlib para visualizar los datos con un diagrama de dispersión.
import matplotlib.pyplot as plt
Consejo: ¿Te apetece aprender distintas formas de visualizar datos en Python? Entonces echa un vistazo al curso Introducción a la visualización de datos con matplotlib.
plt.figure(4, figsize=(10, 8))
plt.scatter(data[:50, 0], data[:50, 1], c='r', label='Iris-setosa')
plt.scatter(data[50:100, 0], data[50:100, 1], c='g',label='Iris-versicolor')
plt.scatter(data[100:, 0], data[100:, 1], c='b',label='Iris-virginica')
plt.xlabel('Sepal length',fontsize=20)
plt.ylabel('Sepal width',fontsize=20)
plt.xticks(fontsize=20)
plt.yticks(fontsize=20)
plt.title('Sepal length vs. Sepal width',fontsize=20)
plt.legend(prop={'size': 18})
plt.show()

En el gráfico anterior, se aprecia claramente una alta correlación para Iris setosa en cuanto a largo y ancho del sépalo. En cambio, hay menos correlación entre Iris versicolor e Iris virginica. Los puntos de versicolor y virginica están más dispersos que los de setosa, que son más densos.
Ahora tracemos también el gráfico para petal-length y petal-width.
plt.figure(4, figsize=(8, 8))
plt.scatter(data[:50, 2], data[:50, 3], c='r', label='Iris-setosa')
plt.scatter(data[50:100, 2], data[50:100, 3], c='g',label='Iris-versicolor')
plt.scatter(data[100:, 2], data[100:, 3], c='b',label='Iris-virginica')
plt.xlabel('Petal length',fontsize=15)
plt.ylabel('Petal width',fontsize=15)
plt.xticks(fontsize=15)
plt.yticks(fontsize=15)
plt.title('Petal length vs. Petal width',fontsize=15)
plt.legend(prop={'size': 20})
plt.show()

También para petal-length y petal-width, el gráfico indica una fuerte correlación en las flores setosa, que aparecen muy agrupadas.
Para validar aún más cómo se correlacionan petal-length y petal-width, tracemos una matriz de correlación para las tres especies.
dataset.iloc[:,2:].corr()
| petal-length | petal-width | |
|---|---|---|
| petal-length | 1.000000 | 0.962865 |
| petal-width | 0.962865 | 1.000000 |
La tabla anterior muestra una fuerte correlación de 0.96 entre petal-length y petal-width cuando se combinan las tres especies.
Analicemos también la correlación entre las especies por separado.
dataset.iloc[:50,:].corr() #setosa
| sepal-length | sepal-width | petal-length | petal-width | |
|---|---|---|---|---|
| sepal-length | 1.000000 | 0.742547 | 0.267176 | 0.278098 |
| sepal-width | 0.742547 | 1.000000 | 0.177700 | 0.232752 |
| petal-length | 0.267176 | 0.177700 | 1.000000 | 0.331630 |
| petal-width | 0.278098 | 0.232752 | 0.331630 | 1.000000 |
dataset.iloc[50:100,:].corr() #versicolor
| sepal-length | sepal-width | petal-length | petal-width | |
|---|---|---|---|---|
| sepal-length | 1.000000 | 0.525911 | 0.754049 | 0.546461 |
| sepal-width | 0.525911 | 1.000000 | 0.560522 | 0.663999 |
| petal-length | 0.754049 | 0.560522 | 1.000000 | 0.786668 |
| petal-width | 0.546461 | 0.663999 | 0.786668 | 1.000000 |
dataset.iloc[100:,:].corr() #virginica
| sepal-length | sepal-width | petal-length | petal-width | |
|---|---|---|---|---|
| sepal-length | 1.000000 | 0.457228 | 0.864225 | 0.281108 |
| sepal-width | 0.457228 | 1.000000 | 0.401045 | 0.537728 |
| petal-length | 0.864225 | 0.401045 | 1.000000 | 0.322108 |
| petal-width | 0.281108 | 0.537728 | 0.322108 | 1.000000 |
De las tres tablas anteriores se desprende que la correlación entre petal-length y petal-width para setosa y virginica es de 0.33 y 0.32 respectivamente, mientras que para versicolor es de 0.78.
Ahora, visualiza la distribución de características trazando los histogramas:
fig = plt.figure(figsize = (8,8))
ax = fig.gca()
dataset.hist(ax=ax)
plt.show()

petal-length, petal-width y sepal-length muestran una distribución unimodal, mientras que sepal-width presenta una distribución de tipo gaussiana. Este análisis es útil para plantearte algoritmos que funcionen bien con este tipo de distribuciones.
A continuación, analizarás si los cuatro atributos están en la misma escala; esto es un aspecto esencial en ML. El data frame de pandas tiene una función incorporada, describe, que te da count, mean, max, min, etc., en formato tabular.
dataset.describe()
| sepal-length | sepal-width | petal-length | petal-width | |
|---|---|---|---|---|
| count | 150.000000 | 150.000000 | 150.000000 | 150.000000 |
| mean | 5.843333 | 3.057333 | 3.758000 | 1.199333 |
| std | 0.828066 | 0.435866 | 1.765298 | 0.762238 |
| min | 4.300000 | 2.000000 | 1.000000 | 0.100000 |
| 25% | 5.100000 | 2.800000 | 1.600000 | 0.300000 |
| 50% | 5.800000 | 3.000000 | 4.350000 | 1.300000 |
| 75% | 6.400000 | 3.300000 | 5.100000 | 1.800000 |
| max | 7.900000 | 4.400000 | 6.900000 | 2.500000 |
Puedes ver que los cuatro atributos están en una escala similar entre 0 y 8, y están en centímetros. Si quieres, puedes reescalarlos entre 0 y 1.
Aunque ya sabemos que hay 50 muestras por clase (aprox. un 33,3% de la distribución total), ¡revisémoslo!
print(dataset.groupby('species').size())
species
Iris-setosa 50
Iris-versicolor 50
Iris-virginica 50
dtype: int64
Preprocesa tus datos
Tras cargar y analizar los datos en detalle, es momento de prepararlos para alimentar tu modelo de ML. En esta sección, los preprocesarás de dos formas: normalización y división en conjuntos de entrenamiento y prueba.
Normaliza tus datos
Hay dos formas habituales de normalizar:
- Normalización por ejemplo, donde normalizas cada muestra individualmente,
- Normalización por característica, en la que normalizas cada característica del mismo modo en todas las muestras.
Entonces, ¿por qué o cuándo necesitas normalizar tus datos? ¿Hace falta estandarizar los datos de Iris?
La respuesta, en general, es que casi siempre es buena práctica. Normalizar pone todas las muestras en una misma escala y rango. Es crucial cuando tus datos no son consistentes. Puedes comprobar la consistencia con la función describe() vista arriba, que te da los valores max y min. Si los valores max y min de una característica son mucho mayores que los de otra, conviene normalizarlas a la misma escala.
Imagina que X es una característica con un rango grande e Y otra con un rango pequeño. Entonces, la influencia de Y puede quedar eclipsada por la de X. En ese caso, es importante normalizar ambas.
En los datos de Iris, la normalización no es necesaria.
Imprimamos de nuevo describe() para ver por qué no hace falta normalizar.
dataset.describe()
| sepal-length | sepal-width | petal-length | petal-width | |
|---|---|---|---|---|
| count | 150.000000 | 150.000000 | 150.000000 | 150.000000 |
| mean | 5.843333 | 3.057333 | 3.758000 | 1.199333 |
| std | 0.828066 | 0.435866 | 1.765298 | 0.762238 |
| min | 4.300000 | 2.000000 | 1.000000 | 0.100000 |
| 25% | 5.100000 | 2.800000 | 1.600000 | 0.300000 |
| 50% | 5.800000 | 3.000000 | 4.350000 | 1.300000 |
| 75% | 6.400000 | 3.300000 | 5.100000 | 1.800000 |
| max | 7.900000 | 4.400000 | 6.900000 | 2.500000 |
El atributo sepal-length va de 4.3 a 7.9; sepal-width, de 2 a 4.4; petal-length, de 1 a 6.9; y petal-width, de 0.1 a 2.5. Todos los valores están entre 0.1 y 7.9, un rango aceptable. Por tanto, no necesitas aplicar normalización al conjunto Iris.
Dividir los datos
Este es otro aspecto clave del aprendizaje automático, ya que tu objetivo es crear un modelo capaz de tomar decisiones o clasificar datos en un entorno de prueba sin intervención humana. Antes de desplegarlo, debes asegurarte de que generaliza bien en datos de test.
Para ello, necesitas conjuntos de entrenamiento y prueba. En Iris, tienes 150 muestras: entrenarás el modelo con el 80% y usarás el 20% restante para probar.
En ciencia de datos verás a menudo el término overfitting, que significa que el modelo aprende demasiado bien los datos de entrenamiento pero falla en los de prueba. Dividir los datos en entrenamiento y prueba (o validación) te ayuda a detectar si tu modelo sobreajusta.
Para la división entrenamiento/prueba, usarás la librería sklearn y su función train_test_split. Vamos a dividir.
from sklearn.model_selection import train_test_split
train_data, test_data, train_label, test_label = train_test_split(dataset.iloc[:,:3], dataset.iloc[:,3], test_size=0.2, random_state=42)
Ten en cuenta que random_state es una semilla: si cambias el número, también cambiará el reparto de los datos. Si mantienes random_state igual y ejecutas varias veces, la división no cambiará.
Imprimamos rápidamente la forma de los datos de entrenamiento y prueba y de sus etiquetas.
train_data.shape,train_label.shape,test_data.shape,test_label.shape
((120, 3), (120,), (30, 3), (30,))
¡Por fin, es hora de alimentar los datos al algoritmo de k-vecinos más cercanos!
El modelo KNN
Tras cargar, analizar y preprocesar, toca introducir los datos en el modelo KNN. Para ello usarás la función neighbors de sklearn, que incluye la clase KNeigborsClassifier.
Empecemos importando el clasificador.
from sklearn.neighbors import KNeighborsClassifier
Nota: el parámetro k (n_neighbors) suele ser un número impar para evitar empates en la votación.
Para decidir el mejor valor del hiperparámetro k, harás una grid-search. Entrenarás y probarás tu modelo con 10 valores distintos de k y, al final, te quedarás con el que mejores resultados dé.
Inicializa una variable neighbors(k) con valores de 1 a 9 y dos matrices de ceros de NumPy, train_accuracy y test_accuracy, para almacenar las precisiones de entrenamiento y prueba. Luego las usarás para trazar un gráfico y elegir el mejor valor de neighbor.
neighbors = np.arange(1,9)
train_accuracy =np.zeros(len(neighbors))
test_accuracy = np.zeros(len(neighbors))
En el siguiente bloque ocurre toda la magia. Harás un enumerate sobre los nueve valores de vecinos y, para cada uno, predecirás tanto en entrenamiento como en prueba. Por último, guardarás la precisión en los arrays de NumPy train_accuracy y test_accuracy.
for i,k in enumerate(neighbors):
knn = KNeighborsClassifier(n_neighbors=k)
#Fit the model
knn.fit(train_data, train_label)
#Compute accuracy on the training set
train_accuracy[i] = knn.score(train_data, train_label)
#Compute accuracy on the test set
test_accuracy[i] = knn.score(test_data, test_label)
A continuación, trazarás con matplotlib las precisiones de entrenamiento y prueba. Con el gráfico de accuracy vs. número de vecinos podrás elegir el mejor k para tu modelo.
plt.figure(figsize=(10,6))
plt.title('KNN accuracy with varying number of neighbors',fontsize=20)
plt.plot(neighbors, test_accuracy, label='Testing Accuracy')
plt.plot(neighbors, train_accuracy, label='Training accuracy')
plt.legend(prop={'size': 20})
plt.xlabel('Number of neighbors',fontsize=20)
plt.ylabel('Accuracy',fontsize=20)
plt.xticks(fontsize=20)
plt.yticks(fontsize=20)
plt.show()

A la vista del gráfico, parece que con n_neighbors=3 el modelo rinde mejor. Así que nos quedamos con n_neighbors=3 y volvemos a entrenar.
knn = KNeighborsClassifier(n_neighbors=3)
#Fit the model
knn.fit(train_data, train_label)
#Compute accuracy on the training set
train_accuracy = knn.score(train_data, train_label)
#Compute accuracy on the test set
test_accuracy = knn.score(test_data, test_label)
Evalúa tu modelo
En el último tramo del tutorial, evaluarás tu modelo en los datos de prueba con un par de técnicas: confusion_matrix y classification_report.
Primero, comprueba la precisión del modelo en el conjunto de prueba.
test_accuracy
0.9666666666666667
¡Voilà! Parece que el modelo clasificó correctamente el 96,66% de los datos de prueba. ¿No está nada mal? Con solo unas pocas líneas de código, has entrenado un modelo de ML que te dice el nombre de la flor usando solo cuatro características, con un 96,66% de acierto. Quién sabe, quizá incluso mejor que una persona.
Matriz de confusión
La matriz de confusión se usa para describir el rendimiento de tu modelo en los datos de prueba, para los que conoces los valores verdaderos o etiquetas.
Scikit-learn ofrece una función que calcula la matriz de confusión por ti.
prediction = knn.predict(test_data)
La siguiente función plot_confusion_matrix() ha sido modificada y tomada de esta fuente.
import itertools
def plot_confusion_matrix(cm, classes,
normalize=False,
title='Confusion matrix',
cmap=plt.cm.Blues):
plt.imshow(cm, interpolation='nearest', cmap=cmap)
plt.title(title)
plt.colorbar()
tick_marks = np.arange(len(classes))
plt.xticks(tick_marks, classes, rotation=45)
plt.yticks(tick_marks, classes)
fmt = '.2f' if normalize else 'd'
thresh = cm.max() / 2.
for i, j in itertools.product(range(cm.shape[0]), range(cm.shape[1])):
plt.text(j, i, format(cm[i, j], fmt),
horizontalalignment="center",
color="white" if cm[i, j] > thresh else "black")
plt.ylabel('True label',fontsize=30)
plt.xlabel('Predicted label',fontsize=30)
plt.tight_layout()
plt.xticks(fontsize=18)
plt.yticks(fontsize=18)
class_names = load_iris().target_names
# Compute confusion matrix
cnf_matrix = confusion_matrix(test_label, prediction)
np.set_printoptions(precision=2)
# Plot non-normalized confusion matrix
plt.figure(figsize=(10,8))
plot_confusion_matrix(cnf_matrix, classes=class_names)
plt.title('Confusion Matrix',fontsize=30)
plt.show()

En la confusion_matrix anterior, se observa que el modelo clasificó todas las flores correctamente salvo una virginica, que se clasificó como versicolor.
Informe de clasificación
El informe de clasificación te ayuda a identificar con más detalle las clases mal clasificadas, mostrando precision, recall y F1 score para cada una. Usarás la librería sklearn para visualizarlo.
from sklearn.metrics import classification_report
print(classification_report(test_label, prediction))
precision recall f1-score support
Iris-setosa 1.00 1.00 1.00 10
Iris-versicolor 0.90 1.00 0.95 9
Iris-virginica 1.00 0.91 0.95 11
micro avg 0.97 0.97 0.97 30
macro avg 0.97 0.97 0.97 30
weighted avg 0.97 0.97 0.97 30
¡Sigue avanzando!
¡Enhorabuena a quienes han llegado hasta el final! Pero esto es solo el principio. ¡Aún queda mucho por descubrir!
Este tutorial se ha centrado en los fundamentos del aprendizaje automático y en la implementación de un tipo de algoritmo, KNN, con Python. El conjunto de datos Iris que usaste es pequeño y relativamente sencillo.
Si este tutorial te ha despertado el interés por saber más, prueba con otros conjuntos de datos o aprende otros algoritmos de ML y aplícalos a Iris para ver cómo afecta a la precisión. Así aprenderás mucho más que con la sola teoría.
Si ya has experimentado lo suficiente con lo básico presentado aquí y con otros algoritmos de ML, quizá quieras profundizar en Python y el análisis de datos.
