Как строить диаграммы рассеивания в Python с помощью plt.scatter()

Разбираем plt.scatter() в Matplotlib. Создаём scatter-графики, настраиваем маркеры по размеру, цвету и форме, используем цветовые карты и фильтруем данные масками.

Обложка: Как строить диаграммы рассеивания в Python с помощью plt.scatter()

Быстрый старт с plt.scatter()

Чтобы начать, установите Matplotlib через pip. После этого достаточно импортировать pyplot под стандартным псевдонимом plt, подготовить два массива данных и вызвать plt.scatter(). В обычном скрипте график нужно явно показать через plt.show(), а в Jupyter Notebook или интерактивной консоли это необязательно.

			import matplotlib.pyplot as plt

price = [2.50, 1.23, 4.02, 3.25, 5.00, 4.40]
sales_per_day = [34, 62, 49, 22, 13, 19]

plt.scatter(price, sales_per_day)
plt.show()
		

Полученный график показывает обратную зависимость: чем выше цена, тем ниже средние продажи. При этом напиток за 4.02 выбивается из общей тенденции, что намекает на его особую популярность и требует дальнейшего анализа.

plt.scatter() против plt.plot()

Тот же базовый график можно построить функцией plt.plot(), передав маркер "o". Однако замеры через timeit обычно показывают, что plt.plot() работает в несколько раз быстрее. Зачем тогда нужен scatter()? Дело в гибкости: только plt.scatter() позволяет задавать индивидуальный размер, цвет и прозрачность для каждой точки.

			import timeit

print("plt.scatter():", timeit.timeit(
    "plt.scatter(price, sales_per_day)",
    number=1000, globals=globals()
))
print("plt.plot():", timeit.timeit(
    "plt.plot(price, sales_per_day, 'o')",
    number=1000, globals=globals()
))
		

Правило простое: для сырых точек без изысков выбирайте plt.plot(), а если нужно раскодировать дополнительные переменные визуально, используйте plt.scatter().

Кастомизация маркеров

Размер точек: параметр s

Допустим, владелец кофейни хочет добавить на график информацию о марже каждого напитка. Параметр s отвечает за площадь маркера. Чтобы разница была заметна, удобно передать массив значений и отмасштабировать его, например, умножив на 10.

			import numpy as np

price = np.asarray([2.50, 1.23, 4.02, 3.25, 5.00, 4.40])
sales_per_day = np.asarray([34, 62, 49, 22, 13, 19])
profit_margin = np.asarray([20, 35, 40, 20, 27.5, 15])

plt.scatter(x=price, y=sales_per_day, s=profit_margin * 10)
plt.show()
		

Цвет: параметр c

Цвет помогает выделить категории. Можно задать его вручную через RGB-кортежи: зелёный для низкого содержания сахара, жёлтый для среднего и красный для высокого. Параметр c принимает список цветов той же длины, что и данные.

			low = (0, 1, 0)
medium = (1, 1, 0)
high = (1, 0, 0)

sugar_content = [low, high, medium, medium, high, low]

plt.scatter(
    x=price,
    y=sales_per_day,
    s=profit_margin * 10,
    c=sugar_content,
)
plt.show()
		

Форма: параметр marker

Когда на одном полотне оказываются два разных набора, важно различать их на глаз. По умолчанию точки рисуются кругами "o", но для второго набора можно выбрать другой символ, например, ромб "d".

			price_cereal = np.asarray([1.50, 2.50, 1.15, 1.95])
sales_cereal = np.asarray([67, 34, 36, 12])
margin_cereal = np.asarray([20, 42.5, 33.3, 18])
sugar_cereal = [low, high, medium, low]

plt.scatter(x=price, y=sales_per_day, s=profit_margin * 10, c=sugar_content)
plt.scatter(
    x=price_cereal,
    y=sales_cereal,
    s=margin_cereal * 10,
    c=sugar_cereal,
    marker="d",
)
plt.show()
		

Прозрачность: параметр alpha

Если точки накладываются друг на друга, часть данных становится невидимой. Параметр alpha задаёт прозрачность от 0 до 1. Значение 0.5 делает маркеры полупрозрачными, и совпадающие наблюдения перестают прятаться.

			plt.scatter(x=price, y=sales_per_day, s=profit_margin * 10, c=sugar_content, alpha=0.5)
plt.scatter(
    x=price_cereal,
    y=sales_cereal,
    s=margin_cereal * 10,
    c=sugar_cereal,
    marker="d",
    alpha=0.5,
)
plt.title("Продажи и цены")
plt.xlabel("Цена")
plt.ylabel("Средние продажи")
plt.legend(["Напитки", "Батончики"])
plt.show()
		

Непрерывный цвет и стили

Вместо фиксированных RGB-цветов можно передать в c числовой массив и выбрать цветовую карту через cmap. Это превращает маркеры в градиент, а plt.colorbar() добавит шкалу значений. Кроме того, оформление можно поменять глобально: plt.style.use("seaborn-v0_8") применит стилистику, близкую к Seaborn.

			sugar_orange = [15, 35, 22, 27, 38, 14]
sugar_cereal = [21, 49, 29, 24]

plt.scatter(
    x=price,
    y=sales_per_day,
    s=profit_margin * 10,
    c=sugar_orange,
    cmap="jet",
    alpha=0.5,
)
plt.scatter(
    x=price_cereal,
    y=sales_cereal,
    s=margin_cereal * 10,
    c=sugar_cereal,
    cmap="jet",
    marker="d",
    alpha=0.5,
)
plt.colorbar()
plt.show()
		

Список доступных стилей возвращает команда plt.style.available. Экспериментируйте с ними, чтобы быстро привести график к единому виду без ручной настройки каждой детали.

Продвинутая техника: маскирование данных

С помощью булевых массивов NumPy можно разбить точки на группы прямо на графике. Представьте, что вы моделируете расписание автобусов: сгенерировали случайные минуты и вероятности, наложили теоретическое распределение, а затем оставили только те точки, что попадают под кривую.

			import random

n_buses = 40
bus_times = np.asarray([random.randint(0, 59) for _ in range(n_buses)])
bus_likelihood = np.asarray([random.random() for _ in range(n_buses)])

x = np.linspace(0, 59, 60)
mean = (15, 45)
sd = (5, 7)
dist = (
    np.exp(-0.5 * ((x - mean[0]) / sd[0]) ** 2)
    + 0.9 * np.exp(-0.5 * ((x - mean[1]) / sd[1]) ** 2)
)
dist = dist / max(dist)

in_region = bus_likelihood < dist[bus_times]
out_region = bus_likelihood >= dist[bus_times]

plt.scatter(x=bus_times[in_region], y=bus_likelihood[in_region], color="green")
plt.scatter(
    x=bus_times[out_region],
    y=bus_likelihood[out_region],
    color="red",
    marker="x",
)
plt.plot(x, dist)
plt.show()
		

Шпаргалка по параметрам

Шпаргалка
Ключевые параметры plt.scatter()
Что за что отвечает в функции
  • x, y — координаты точек (обязательные аргументы).
  • s — размер маркера: одно число или массив значений для каждой точки.
  • c — цвет в формате RGB, название или массив чисел для цветовой карты.
  • marker — форма точки: 'o' (круг), 'd' (ромб), 'x' (крест) и другие.
  • cmap — название цветовой карты, применяется вместе с числовым c.
  • alpha — прозрачность от 0 (полностью прозрачный) до 1 (непрозрачный).

Выводы

Функция plt.scatter() превращает плоскую диаграмму рассеяния в многослойный инструмент анализа. На одном полотне можно одновременно показать цену, объём продаж, прибыльность, категорию товара и его характеристики через размер, цвет и форму маркеров. Для простых задач по-прежнему удобен plt.plot(), но стоит ли задача исследовать сложные взаимосвязи, scatter() становится незаменим.