Weiter zum Inhalt

Einführung in Machine Learning mit Python

In diesem Tutorial lernst du die Welt des Machine Learnings (ML) mit Python kennen. Um ML praktisch zu verstehen, verwendest du den bekannten Algorithmus K-Nearest Neighbor (KNN) in Python.
Aktualisiert 18. Sept. 2026  · 14 Min. lesen

Mit KI erkunden

ChatGPTClaudePerplexity

Du implementierst KNN auf dem berühmten Iris-Datensatz.

Hinweis: Vielleicht möchtest du zuvor den Kurs Machine Learning mit Python belegen. Für Hintergründe zur Entwicklung von ML und vieles mehr lies diesen Beitrag.

Einführung

Machine Learning ist aus der Informatik hervorgegangen und befasst sich mit der Entwicklung von Algorithmen, die aus Erfahrung lernen. Dafür brauchen sie Daten mit bestimmten Attributen, anhand derer die Algorithmen sinnvolle, vorhersagbare Muster finden. ML-Aufgaben lassen sich grob in Concept Learning, Clustering, Predictive Modeling usw. einteilen. Das Ziel von ML-Algorithmen ist es, Entscheidungen korrekt und ohne menschliches Eingreifen treffen zu können. Aktienkurse oder Wetter vorhersagen sind zwei typische Anwendungsfälle.

Es gibt zahlreiche ML-Algorithmen wie Decision Trees, Naive Bayes, Random Forest, Support Vector Machine, K-Nearest Neighbor, K-Means Clustering usw.

In diesem Tutorial verwendest du aus dieser Familie den Algorithmus k-Nearest Neighbor.

Was genau steckt nun hinter dem K-Nearest-Neighbor-Algorithmus? Schauen wir es uns an!

Was ist k-Nearest Neighbor?

KNN, also k-Nearest Neighbor, ist ein überwachter Lernalgorithmus. „Überwacht“ bedeutet hier, dass während der Lernphase die Klassenlabels der Trainingsdaten verwendet werden. Es ist ein instanzbasierter ML-Algorithmus: Neue Datenpunkte werden auf Basis gespeicherter, gelabelter Instanzen klassifiziert. KNN eignet sich für Klassifikation und Regression, wird aber häufiger für Klassifikation eingesetzt.

Das k in KNN ist eine entscheidende Variable (Hyperparameter), die hilft, einen Datenpunkt korrekt zu klassifizieren. Genauer gesagt ist k die Anzahl der nächsten Nachbarn, deren Stimmen du bei der Klassifikation eines neuen Datenpunkts berücksichtigst.

visualization of knn visualization of knn

Abbildung 1. Visualisierung von KNN Quelle

Du siehst: Wenn der Wert k von 1 auf 7 steigt, wird die Entscheidungsgrenze zwischen zwei Klassen mit einigen Datenpunkten glatter.

Wie funktioniert die „Magie“, dass ein neuer Datenpunkt jedes Mal anhand der gespeicherten Punkte korrekt zugeordnet wird?

Schnell erklärt, in diesen Schritten:

  • Zuerst lädst du alle Daten und legst einen Wert für k fest,
  • dann berechnest du den Abstand zwischen den gespeicherten Datenpunkten und dem neuen Punkt, den du klassifizieren willst — z. B. mit Manhattan-Distanz (L1), Euklidischer Distanz (L2), Kosinus-Ähnlichkeit, Bhattacharyya-Distanz, Chebyshev-Distanz usw.
  • Anschließend sortierst du die Abstände auf- oder absteigend und bestimmst die k nächsten Nachbarn.
  • Du sammelst die Labels dieser k Nachbarn und klassifizierst den neuen Punkt per Mehrheitsvotum oder gewichteter Abstimmung. Das Label des Datenpunkts mit der höchsten Wertung setzt sich durch.
  • Zum Schluss gibst du die vorhergesagte Klasse für die neue Instanz zurück.

Vorhersagen gibt es in zwei Varianten: Klassifikation (ein Klassenlabel wird zugewiesen) oder Regression (es wird ein Wert zugewiesen). Bei Regression erhält der neue Datenpunkt typischerweise den Mittelwert der k nächsten Nachbarn.

Nachteile von KNN: Erstens ist die Suche nach den nächsten Nachbarn für jeden neuen Datenpunkt rechenintensiv. Zweitens ist die Wahl eines guten k-Werts oft mühsam. Drittens ist nicht immer klar, welche Distanzmetrik sich am besten eignet.

Genug Theorie, oder? Lass uns die Daten laden, analysieren und verstehen, mit denen du heute arbeitest.

Iris-Daten laden

Der Iris-Datensatz besteht aus 150 Stichproben mit drei Klassen: Iris-Setosa, Iris-Versicolor und Iris-Virginica. Vier Merkmale/Attribute dienen zur eindeutigen Zuordnung zu einer der drei Klassen: sepal-length, sepal-width, petal-length und petal-width.

Du kannst gern auch einen anderen öffentlichen oder deinen eigenen Datensatz verwenden.

Sklearn ist eine weit verbreitete Machine-Learning-Bibliothek für Python, die viele Aufgaben in der Data Science abdeckt. Sie enthält diverse Algorithmen für Klassifikation, Regression und Clustering, darunter support vector machines, random forests, gradient boosting, k-means, KNN usw. In sklearn gibt es das Modul datasets mit mehreren gebrauchsfertigen Datensätzen, darunter Iris. Das Laden ist intuitiv und unkompliziert. Also, laden wir den iris-Datensatz.

from sklearn.datasets import load_iris

load_iris liefert sowohl die Daten als auch die Klassenlabels für jede Stichprobe. Lass uns das schnell extrahieren.

data = load_iris().data

Die Variable data ist ein NumPy-Array der Form (150,4) mit 150 Stichproben und je vier Attributen. Jede Klasse hat 50 Stichproben.

data.shape
(150, 4)

Als Nächstes holen wir die Klassenlabels.

labels = load_iris().target
labels.shape
(150,)

Nun musst du Daten und Labels kombinieren. Dafür verwendest du die hervorragende Python-Bibliothek NumPy. NumPy unterstützt große, mehrdimensionale Arrays und Matrizen und liefert eine umfangreiche Sammlung mathematischer Funktionen, um mit diesen Arrays zu arbeiten. Also importieren wir sie schnell!

import numpy as np

Da data ein 2D-Array ist, solltest du auch labels in ein 2D-Array umformen.

labels = np.reshape(labels,(150,1))

Jetzt nutzt du die concatenate-Funktion aus numpy und setzt axis=-1, um entlang der zweiten Dimension zu konkatenieren.

data = np.concatenate([data,labels],axis=-1)
data.shape
(150, 5)

Als Nächstes importierst du die Datenanalysebibliothek pandas. Sie eignet sich hervorragend, um Daten tabellarisch anzuordnen und Operationen/Transformationen darauf durchzuführen. Insbesondere bietet sie Datenstrukturen und Funktionen zur Arbeit mit numerischen Tabellen und Zeitreihen.

In diesem Tutorial wirst du pandas intensiv verwenden.

import pandas as pd
names = ['sepal-length', 'sepal-width', 'petal-length', 'petal-width', 'species']
dataset = pd.DataFrame(data,columns=names)

Jetzt hast du das DataFrame dataset mit Daten und Klassenlabels, die du brauchst!

Bevor wir weitermachen: Die Variable labels enthält die Klassenlabels als Zahlen. Wir wandeln sie gleich in die entsprechenden Blumennamen (Species) um.

Dazu wählst du die Spalte class aus und ersetzt die drei numerischen Werte durch die Species. Mit inplace=True änderst du das DataFrame dataset direkt.

dataset['species'].replace(0, 'Iris-setosa',inplace=True)
dataset['species'].replace(1, 'Iris-versicolor',inplace=True)
dataset['species'].replace(2, 'Iris-virginica',inplace=True)

Lass uns die ersten fünf Zeilen von dataset ausgeben und anschauen!

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

Daten analysieren

Schauen wir uns schnell an, wie die drei Blumenarten aussehen und worin sie sich unterscheiden — nicht nur numerisch, sondern auch in der Realität!

iris(Quelle)

Visualisieren wir jetzt die oben geladenen Daten mit einem scatterplot, um zu sehen, wie stark sich zwei Variablen gegenseitig beeinflussen — sprich, wie hoch ihre Korrelation ist.

Dafür verwendest du die Bibliothek matplotlib.

import matplotlib.pyplot as plt

Tipp: Du willst mehr über Datenvisualisierung in Python lernen? Dann schau dir den Kurs Introduction to data visualization with matplotlib an.

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()
sepal length x width scatter plot

Aus dem Plot ist klar zu erkennen: Bei Iris setosa besteht in Bezug auf Kelchblattlänge und -breite eine starke Korrelation. Bei Iris versicolor und Iris virginica ist die Korrelation geringer. Die Datenpunkte von versicolor und virginica sind stärker gestreut, während setosa dichter beisammenliegt.

Plotten wir jetzt noch den Graphen für petal-length und 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()
petal length x width scatter plot

Auch hier, bei petal-length und petal-width, zeigt sich für setosa eine starke Korrelation mit dicht beieinanderliegenden Punkten.

Zur weiteren Untermauerung der Korrelation zwischen petal-length und petal-width plotten wir eine Korrelationsmatrix über alle drei Arten.

dataset.iloc[:,2:].corr()
  petal-length petal-width
petal-length 1.000000 0.962865
petal-width 0.962865 1.000000

Die Tabelle zeigt eine starke Korrelation von 0.96 zwischen petal-length und petal-width, wenn alle drei Arten kombiniert werden.

Betrachten wir die Korrelationen auch getrennt nach den drei Arten.

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

Aus den drei Tabellen wird deutlich: Die Korrelation zwischen petal-length und petal-width beträgt bei setosa 0.33 und bei virginica 0.32. Für versicolor liegt sie bei 0.78.

Als Nächstes visualisieren wir die Merkmalsverteilung mit Histogrammen:

fig = plt.figure(figsize = (8,8))
ax = fig.gca()
dataset.hist(ax=ax)
plt.show()
bar charts

petal-length, petal-width und sepal-length zeigen eine unimodale Verteilung, während sepal-width eher einer Gauß-Verteilung ähnelt. Diese Analysen sind hilfreich, um Algorithmen zu wählen, die mit solchen Verteilungen gut arbeiten.

Nun prüfen wir, ob alle vier Attribute auf derselben Skala liegen — ein wichtiger Aspekt in ML. Das pandas-DataFrame hat mit describe eine Funktion, die dir count, mean, max, min usw. tabellarisch liefert.

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

Du siehst: Alle vier Attribute liegen in ähnlichen Skalen zwischen 0 und 8 und sind in Zentimetern. Falls gewünscht, könntest du zusätzlich auf den Bereich 0 bis 1 skalieren.

Auch wenn wir wissen, dass es 50 Stichproben pro Klasse gibt (~33,3 % der Gesamtverteilung), prüfen wir es kurz nach.

print(dataset.groupby('species').size())
species
Iris-setosa        50
Iris-versicolor    50
Iris-virginica     50
dtype: int64

Daten vorbereiten

Nach dem Laden und der Analyse bereitest du die Daten für dein ML-Modell auf. In diesem Abschnitt normalisierst du die Daten (falls nötig) und teilst sie in Trainings- und Testdaten.

Daten normalisieren

Es gibt zwei gängige Wege, Daten zu normalisieren:

  • Beispiel-Normalisierung (jede Stichprobe einzeln normalisieren),
  • Feature-Normalisierung (jedes Merkmal über alle Stichproben gleich normalisieren).

Warum und wann normalisieren? Und muss der Iris-Datensatz standardisiert werden?

Im Grunde fast immer ist Normalisierung eine gute Idee, weil sie alle Stichproben auf dieselbe Skala bringt. Besonders wichtig ist sie bei inkonsistenten Daten. Mit describe() kannst du max und min prüfen. Sind die Wertebereiche zweier Features stark unterschiedlich, sollte man beide auf dieselbe Skala bringen.

Hat Feature X einen viel größeren Bereich als Feature Y, kann X den Einfluss von Y überdecken. Dann ist eine Normalisierung beider Features sinnvoll.

Beim Iris-Datensatz ist Normalisierung nicht erforderlich.

Wirf noch einmal einen Blick auf describe(), um zu sehen, warum hier keine Normalisierung nötig ist.

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

sepal-length reicht von 4,3 bis 7,9, sepal-width von 2 bis 4,4, petal-length von 1 bis 6,9 und petal-width von 0,1 bis 2,5. Alle Features liegen damit zwischen 0,1 und 7,9 — das ist akzeptabel. Du musst den Iris-Datensatz daher nicht normalisieren.

Daten aufteilen

Ein weiterer Kernschritt im Machine Learning: Dein Modell soll in einer Testsituation ohne menschliches Eingreifen korrekt entscheiden bzw. klassifizieren. Bevor du es produktiv einsetzt, musst du sicherstellen, dass es auf Testdaten gut generalisiert.

Dafür brauchst du Trainings- und Testdaten. Beim Iris-Datensatz mit 150 Stichproben trainierst du auf 80 % der Daten und testest auf den verbleibenden 20 %.

In der Data Science begegnest du oft dem Begriff Overfitting: Das Modell lernt die Trainingsdaten zu gut, performt aber schlecht auf Testdaten. Die Aufteilung in Training und Test (oder Validierung) hilft dir, Overfitting zu erkennen.

Für das Splitten nutzt du sklearn mit der Funktion train_test_split. Los geht’s.

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) 

Beachte: random_state ist ein Seed. Änderst du ihn, ändert sich auch die Aufteilung. Hältst du ihn konstant und führst die Zelle mehrfach aus, bleibt der Split identisch.

Drucken wir kurz die Shapes der Trainings- und Testdaten samt Labels.

train_data.shape,train_label.shape,test_data.shape,test_label.shape
((120, 3), (120,), (30, 3), (30,))

Jetzt füttern wir die Daten in den k-Nearest-Neighbor-Algorithmus!

Das KNN-Modell

Nach Laden, Analyse und Vorbereitung gibst du die Daten nun in das KNN-Modell. Dafür verwendest du in sklearn das Modul neighbors mit der Klasse KNeighborsClassifier.

Importieren wir zuerst den Classifier.

from sklearn.neighbors import KNeighborsClassifier

Hinweis: Der Parameter k (n_neighbors) ist oft ungerade, um Stimmengleichstand zu vermeiden.

Um den besten Wert für den Hyperparameter k zu finden, nutzt du eine Grid-Search. Du trainierst und testest das Modell mit 10 unterschiedlichen k-Werten und nimmst am Ende den besten.

Dazu initialisieren wir neighbors(k) mit Werten von 1–9 sowie zwei NumPy-Null-Arrays train_accuracy und test_accuracy zur Speicherung der Trainings- und Testgenauigkeit. Sie brauchst du gleich für den Auswahlplot.

neighbors = np.arange(1,9)
train_accuracy =np.zeros(len(neighbors))
test_accuracy = np.zeros(len(neighbors))

Im nächsten Codeschnipsel passiert die eigentliche Arbeit: Du enumeratest über alle Nachbarwerte, trainierst jeweils und misst die Genauigkeit auf Trainings- und Testdaten. Die Werte landen in den Arrays train_accuracy und 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)

Anschließend plottest du Trainings- und Testgenauigkeit mit matplotlib. Aus dem Diagramm Accuracy vs. Anzahl Nachbarn wählst du dann das beste k.

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()
knn accuracy chart

Der Plot zeigt: Bei n_neighbors=3 performt das Modell am besten. Also bleiben wir bei n_neighbors=3 und trainieren erneut.

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)

Modell evaluieren

Zum Abschluss bewertest du dein Modell auf den Testdaten mit confusion_matrix und classification_report.

Prüfen wir zuerst die Genauigkeit auf den Testdaten.

test_accuracy
0.9666666666666667

Großartig! Offenbar klassifiziert das Modell 96,66 % der Testdaten korrekt. Nicht schlecht — mit wenigen Zeilen Code hast du ein ML-Modell trainiert, das anhand von nur vier Features den Blumennamen mit 96,66 % Genauigkeit vorhersagt. Vielleicht sogar besser als ein Mensch.

Confusion Matrix

Eine Confusion Matrix beschreibt die Leistung deines Modells auf Testdaten mit bekannten wahren Labels.

Scikit-learn stellt eine Funktion bereit, die die Confusion Matrix für dich berechnet.

prediction = knn.predict(test_data)

Die folgende Funktion plot_confusion_matrix() wurde angepasst und aus dieser Quelle übernommen.

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()
confusion matrix

Aus der confusion_matrix siehst du: Alle Blumen wurden korrekt klassifiziert — bis auf eine virginica, die als versicolor eingestuft wurde.

Classification Report

Der Classification Report zeigt dir fehlerklassifizierte Klassen detaillierter — mit Precision, Recall und F1-Score je Klasse. Zur Darstellung nutzt du die sklearn-Bibliothek.

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

Geh den nächsten Schritt!

Glückwunsch an alle, die bis hierher gekommen sind! Aber das war nur der Anfang. Da geht noch viel mehr!

Dieses Tutorial hat vor allem die Grundlagen von Machine Learning behandelt und einen ML-Algorithmus — KNN — mit Python umgesetzt. Der verwendete Iris-Datensatz ist recht klein und überschaubar.

Wenn dieses Tutorial deine Neugier geweckt hat, probiere weitere Datensätze aus oder lerne andere ML-Algorithmen kennen und wende sie auf den Iris-Datensatz an, um den Einfluss auf die Genauigkeit zu sehen. So lernst du weit mehr als nur die Theorie!

Wenn du die hier gezeigten Grundlagen und weitere ML-Algorithmen durchgespielt hast, lohnt sich der nächste Schritt in Python und Datenanalyse.

Themen
Maschinelles Lernen
Python

Lerne mehr über Machine Learning und Python

Kurs

Machine Learning verstehen

2 Std.
308K
In diesem Kurs lernst du das spannende Themenfeld des maschinellen Lernens kennen – und du benötigst dafür gar keine Programmierkenntnisse.
Details anzeigenRight Arrow
Kurs Starten
Mehr anzeigenRight Arrow