Ошибка при построении диаграммы рассеивания: KeyError: 'class'
Мой код:
import numpy as np
import pandas as pd
import matplotlib.pyplot as plt
%matplotlib
url = 'https://archive.ics.uci.edu/ml/machine-learning-databases/wine-quality/winequality-red.csv'
data = pd.read_csv(url,sep=';')
col1 = 'fixed acidity'
col2 = 'volatile acidity'
plt.figure(figsize=(10, 6))
plt.scatter(data[col1][data['class'] == '+'],
data[col2][data['class'] == '+'],
alpha=0.75,
color='red',
label='+')
plt.scatter(data[col1][data['class'] == '-'],
data[col2][data['class'] == '-'],
alpha=0.75,
color='blue',
label='-')
plt.xlabel(col1)
plt.ylabel(col2)
plt.legend(loc='best');
При выполнении кода выбрасывается исключение KeyError: 'class'.
В чём ошибка? Что нужно поставить на место 'class'?
Ответы (2 шт):
Создается стойкое впечатление, что вы взяли какой-то код, даже не дали себе труд разобраться что он делает, запулили в него данные, получили ошибку и вместо того, что-бе подумать - побежали сюда за помощью. Что-то SW все больше превращается в сайт для решения домашних заданий нерадивых школьников.
Ну ответе, откуда у вас в коде взялся 'class' да еще с плюсиками и минусиками в придачу. Мы не понимаем, что вы там вроде-как надумали делать, но базовый код у вас абсолютно рабочий:
import numpy as np
import pandas as pd
import matplotlib.pyplot as plt
url = 'https://archive.ics.uci.edu/ml/machine-learning-databases/wine-quality/winequality-red.csv'
data = pd.read_csv(url,sep=';')
col1 = 'fixed acidity'
col2 = 'volatile acidity'
plt.figure(figsize=(10, 6))
plt.scatter(data[col1],
data[col2],
alpha=0.75,
color='blue',
label='-')
plt.xlabel(col1)
plt.ylabel(col2)
plt.legend(loc='best');
Результат:
А дальше - думайте сами.
Осмелюсь предположить, что вы хотели нарисовать как соотносятся два указанных признака для белых и красных сортов вина:
import matplotlib
matplotlib.style.use('ggplot')
url = "https://archive.ics.uci.edu/ml/machine-learning-databases/wine-quality/winequality-{}.csv"
red = pd.read_csv(url.format("red"), sep=";")
white = pd.read_csv(url.format("white"), sep=";")
col1 = 'fixed acidity'
col2 = 'volatile acidity'
ax = red.plot.scatter(
x=col1, y=col2, c="purple", edgecolors="silver",
figsize=(10, 10), s=30, label="red", grid=True)
white.plot.scatter(
x=col1, y=col2, c="beige", ax=ax, edgecolors="silver",
s=30, label="white")
результат:

