Агрегации: groupBy, agg, rollup, cube и grouping sets

Механика агрегации в Spark: от shuffle и partial aggregation до rollup/cube/grouping sets. HashAggregate vs SortAggregate, skew в группировке, AQE и построение аналитических витрин.

core

Почему агрегации - сердце аналитики

Любая BI-витрина, OLAP-отчёт или Data Mart - это агрегация. Выручка по регионам, количество уникальных пользователей по дням, средний чек по категориям товаров. Агрегации превращают миллиарды сырых событий в десятки тысяч строк, с которыми работает аналитик.

В Spark агрегация - не просто синтаксический сахар над SQL. Под groupBy().agg() скрывается сложная машинерия: многофазное выполнение, shuffle через сеть, хэш-таблицы в оперативной памяти и кодогенерация через Tungsten. Понять эту механику - значит уметь проектировать агрегации, которые не падают с OOM и не создают straggler-задачи на перекошенных ключах.

Как Spark выполняет агрегацию: от кода к физическому плану

Когда вы пишете df.groupBy("region").agg(F.sum("amount")), Catalyst строит логический план и затем транслирует его в физический. Физическая агрегация в Spark - это всегда минимум два этапа:

Partial Aggregation: почему агрегация быстрее чем кажется

Stage 1 (Map-side): Каждый Executor читает свои локальные партиции и сразу агрегирует данные в памяти. Это называется partial aggregation или map-side combine. Для sum(amount) это означает: вместо того, чтобы передавать по сети все строки с region=MSK, Executor суммирует их локально и передаёт одно число.

Эффект огромный: если в исходных данных 10 миллионов строк с region=MSK, через сеть пойдёт только одно число на каждый Executor.

Stage 2 (Reduce-side): После shuffle каждая выходная партиция содержит частичные агрегаты для одного набора ключей. Финальная агрегация объединяет их.

Почему groupBy - это wide transformation

groupBy - широкая трансформация (wide transformation): она требует shuffle, то есть пересылки данных между Executor'ами. Это создаёт границу Stage в DAG. В отличие от filter и select (narrow transformations), одна строка в результирующей партиции может зависеть от строк из любой входной партиции.

Следствие: каждый groupBy в коде создаёт как минимум один Exchange (shuffle) в физическом плане. Минимизация числа groupBy - прямой путь к ускорению pipeline.

HashAggregate vs SortAggregate

Catalyst выбирает между двумя физическими стратегиями:

HashAggregate (используется в большинстве случаев): строит хэш-таблицу в Tungsten off-heap памяти. Каждая новая строка - O(1) операция поиска ключа и обновления аккумулятора. При нехватке памяти - spill на диск.

SortAggregate (fallback): используется, когда функция агрегации не поддерживает буферизацию через Tungsten (например, некоторые UDAF) или тип данных не сериализуется. Сначала сортирует данные по ключу, затем проходит последовательно. Медленнее, но гарантированно работает при любом объёме.

В explain() вы увидите HashAggregate или SortAggregate в Physical Plan.

Базовый groupBy и agg

Синтаксис

from pyspark.sql import SparkSession
import pyspark.sql.functions as F
from pyspark.sql.functions import col

spark = SparkSession.builder.appName("aggregations").getOrCreate()

orders = spark.read.parquet("/data/orders/")

# Простейший groupBy с одной функцией
orders.groupBy("region").sum("amount").show()

# Предпочтительный стиль: groupBy + agg со списком выражений
result = orders.groupBy("region").agg(
    F.sum("amount").alias("total_revenue"),
    F.count("*").alias("order_count"),
    F.avg("amount").alias("avg_order_value"),
    F.max("amount").alias("max_order"),
)

Почему agg() лучше чем цепочка .sum().count():

Метод .agg() компилируется в один Aggregate узел в физическом плане - одна хэш-таблица, один проход по данным, один shuffle. Цепочка .sum("amount").count("*") (технически невозможна напрямую, но аналог - два отдельных groupBy) создала бы два shuffle. Всегда объединяйте все агрегаты одной группировки в один agg().

Все стандартные функции агрегации

# Подсчёт
F.count("*")                    # все строки, включая NULL
F.count("amount")               # строки где amount IS NOT NULL
F.countDistinct("customer_id")  # уникальные значения (дорого!)
F.approx_count_distinct("customer_id", rsd=0.05)  # ±5% погрешность, быстро

# Суммарные метрики
F.sum("amount")
F.avg("amount")                 # mean
F.mean("amount")                # alias для avg

# Экстремальные значения
F.min("amount")
F.max("amount")
F.min_by("order_id", "amount")  # order_id строки с минимальным amount
F.max_by("order_id", "amount")  # Spark 3.0+

# Статистика
F.stddev("amount")              # sample standard deviation
F.stddev_pop("amount")          # population standard deviation
F.variance("amount")
F.skewness("amount")
F.kurtosis("amount")

# Перцентили (точные и приближённые)
F.percentile_approx("amount", 0.5)           # медиана, ~1% погрешность
F.percentile_approx("amount", [0.25, 0.75])  # Q1 и Q3
F.percentile("amount", 0.95)                 # точный (требует сортировки)

# Коллекции
F.collect_list("product_id")    # список всех значений (с дубликатами)
F.collect_set("product_id")     # уникальные значения (порядок не гарантирован)

# Первое/последнее значение в группе
F.first("status", ignorenulls=True)
F.last("status",  ignorenulls=True)

NULL в агрегации: важные нюансы

Spark обрабатывает NULL в агрегатах по правилам SQL:

Особенность группировки по NULL: если ключ группировки содержит NULL, Spark создаёт для таких строк отдельную группу (в SQL это тоже так). NULL-группа не теряется - она появляется в результате с NULL в колонке ключа.

# Данные: [("MSK", 100), ("SPB", 200), (None, 50), (None, 75)]
df.groupBy("region").agg(F.sum("amount")).show()
# +------+------+
# |region|  sum |
# +------+------+
# |   MSK|   100|
# |   SPB|   200|
# |  null|   125|   ← NULL образует свою группу
# +------+------+

GroupBy по нескольким колонкам

# Двумерный срез: регион × категория
result = orders.groupBy("region", "category").agg(
    F.sum("amount").alias("revenue"),
    F.count("*").alias("orders"),
    F.countDistinct("customer_id").alias("customers"),
)

# Месяц + категория: временной срез
result = orders.groupBy(
    F.year("created_at").alias("year"),
    F.month("created_at").alias("month"),
    "category",
).agg(
    F.sum("amount").alias("monthly_category_revenue"),
)

Alias в агрегациях: обязателен для downstream

Результат агрегации без alias получает автоматическое имя вида sum(amount), count(1). Это создаёт проблемы при последующих операциях и join:

# Плохо: колонка называется "sum(amount)" - неудобно и хрупко
orders.groupBy("region").agg(F.sum("amount"))

# Хорошо: читаемое имя, с которым работают downstream
orders.groupBy("region").agg(
    F.sum("amount").alias("total_revenue"),
    F.count("*").alias("order_count"),
)

Сложные агрегации

collect_list и collect_set: данные в коллекцию

collect_list собирает все значения группы в массив. Это мощный инструмент, но несёт риски:

# Список всех product_id в каждом заказе
order_items = orders.groupBy("order_id").agg(
    F.collect_list("product_id").alias("products"),
    F.collect_list("amount").alias("amounts"),
)

# collect_set: только уникальные значения
unique_categories = orders.groupBy("customer_id").agg(
    F.collect_set("category").alias("purchased_categories"),
)

Проблема: если в одной группе миллион строк, collect_list создаст массив из миллиона элементов. Такой массив хранится в памяти Executor'а. При большом cardinality или несбалансированных группах - OOM.

Когда безопасно: когда количество элементов в каждой группе ограничено (например, позиции в заказе - обычно < 100).

countDistinct vs approx_count_distinct

countDistinct - точный подсчёт уникальных значений. Для этого Spark собирает все значения ключа в одну партицию (shuffle!) и считает уникальные. При высоком cardinality (миллионы уникальных значений) это требует большого объёма shuffle и памяти.

# Точный - требует полного shuffle всех значений
exact = orders.groupBy("region").agg(
    F.countDistinct("customer_id").alias("unique_customers")
)

# Приближённый - HyperLogLog алгоритм, ±rsd погрешность
approx = orders.groupBy("region").agg(
    F.approx_count_distinct("customer_id", rsd=0.05).alias("approx_customers")
)

approx_count_distinct использует алгоритм HyperLogLog: он хранит вероятностную структуру данных вместо полного множества. Типичная погрешность 1–5%. Для dashboardов и мониторинга этого достаточно, но если нужна точная финансовая отчётность - только countDistinct.

Метрика countDistinct approx_count_distinct(rsd=0.05)
Точность 100% ±5%
Shuffle Полный (все значения) Только скетч ~1 KB
Скорость Медленно В 5–20 раз быстрее
Применение Финансы, аудит Аналитика, дашборды

Rollup: иерархические итоги

rollup - это расширение groupBy, которое автоматически генерирует промежуточные итоги по иерархии группировочных колонок.

Если groupBy("year", "month", "region") даёт строки только для каждой тройки значений, то rollup("year", "month", "region") добавляет:

  • итоги по (year, month) - без детализации по region
  • итоги по (year) - без детализации по month и region
  • итог () - общий итог

Rollup следует правостороннему принципу: сворачивает иерархию справа налево. Первая колонка - самый верхний уровень иерархии, последняя - самый детальный.

# Иерархия: год → месяц → регион
report = orders.rollup(
    F.year("created_at").alias("year"),
    F.month("created_at").alias("month"),
    "region",
).agg(
    F.sum("amount").alias("revenue"),
    F.count("*").alias("orders"),
)

report.orderBy("year", "month", "region").show(30)
# +----+-----+------+---------+------+
# |year|month|region|  revenue|orders|
# +----+-----+------+---------+------+
# |2024|    1|   EKB|  12500.0|    45|
# |2024|    1|   MSK|  85000.0|   312|
# |2024|    1|   SPB|  43000.0|   178|
# |2024|    1|  null| 140500.0|   535|  ← итог по месяцу
# |2024|    2|   MSK|  92000.0|   341|
# |2024|    2|  null|  92000.0|   341|  ← итог по месяцу
# |2024| null|  null| 232500.0|   876|  ← итог по году
# |null| null|  null| 232500.0|   876|  ← общий итог
# +----+-----+------+---------+------+

Cube: все комбинации измерений

cube генерирует агрегаты для всех возможных комбинаций группировочных колонок. Для N колонок это 2^N комбинаций.

Для cube("year", "region", "category") будут сгенерированы итоги по:

  • каждой тройке (year, region, category)
  • каждой паре: (year, region), (year, category), (region, category)
  • каждой колонке отдельно: (year), (region), (category)
  • общий итог ()
# Cube: все комбинации year × region × category
cube_result = orders.cube("year", "region", "category").agg(
    F.sum("amount").alias("revenue"),
)

Предупреждение о cardinality explosion: при 10 измерениях cube генерирует 2^10 = 1024 комбинации. Если каждое измерение имеет высокий cardinality (например, тысячи уникальных значений), объём результата взрывается. Cube применяйте только к измерениям с низким cardinality (регион, категория, год, квартал).

Rollup vs Cube: когда что выбрать

Критерий rollup cube
Иерархия данных Чёткая: год → месяц → день Независимые измерения
Число комбинаций N+1 2^N
Применение Временные ряды, geo-иерархия Многомерный OLAP-анализ
Риск взрыва данных Низкий Высокий при N>5

GROUPING и GROUPING_ID: различение строк итогов

После rollup/cube в результате появляются NULL-значения в колонках. Проблема: невозможно отличить «итоговую» строку от строки, где ключ действительно был NULL в исходных данных.

grouping(col) возвращает 1, если колонка участвует в обобщении (строка итога), и 0 иначе.

grouping_id() возвращает число-битовую маску по всем группировочным колонкам.

from pyspark.sql.functions import grouping, grouping_id

report = orders.rollup("year", "month", "region").agg(
    F.sum("amount").alias("revenue"),
    grouping("region").alias("is_region_total"),
    grouping("month").alias("is_month_total"),
    grouping_id("year", "month", "region").alias("level"),
)

# level=0: строка детализации (year, month, region все заданы)
# level=1: итог по (year, month) - region обобщён
# level=3: итог по (year) - month и region обобщены
# level=7: общий итог - все три колонки обобщены

# Фильтрация только строк итогов по месяцу
monthly_totals = report.filter(col("level") == 1)

# Добавить читаемую метку строки
labeled = report.withColumn(
    "row_label",
    F.when(col("level") == 0, F.concat_ws("/", col("year"), col("month"), col("region")))
     .when(col("level") == 1, F.concat_ws("/", col("year"), col("month"), F.lit("ИТОГО")))
     .when(col("level") == 3, F.concat(col("year").cast("string"), F.lit(" ИТОГО")))
     .otherwise(F.lit("ОБЩИЙ ИТОГО"))
)

Grouping Sets: точечный контроль агрегации

GROUPING SETS - SQL-конструкция, которая позволяет явно задать список комбинаций для агрегации. Это более гибкий инструмент, чем rollup или cube: вы выбираете именно нужные комбинации.

В PySpark нет прямого метода .groupingSets() в DataFrame API (только в Spark 3.5+ добавили grouping_sets()), но всегда доступен через spark.sql():

# Зарегистрируем DataFrame как временную таблицу
orders.createOrReplaceTempView("orders")

# Grouping sets: только (region), (year, month) и общий итог
result = spark.sql("""
    SELECT
        region,
        year,
        month,
        SUM(amount)  AS revenue,
        COUNT(*)     AS orders,
        GROUPING_ID(region, year, month) AS level
    FROM orders
    GROUP BY GROUPING SETS (
        (region),
        (year, month),
        ()
    )
""")

Ключевое преимущество GROUPING SETS перед несколькими запросами: Spark сканирует данные один раз. Эквивалент через UNION ALL потребовал бы трёх отдельных scan и трёх shuffle:

# Антипаттерн: три separate scan и три shuffle
by_region = orders.groupBy("region").agg(F.sum("amount"))
by_month  = orders.groupBy("year", "month").agg(F.sum("amount"))
total     = orders.agg(F.sum("amount"))
result    = by_region.unionAll(by_month).unionAll(total)   # три scan!

# Правильно: один scan через GROUPING SETS
result = spark.sql("SELECT ... GROUP BY GROUPING SETS ((region), (year, month), ())")

В Spark 3.5+ появился метод grouping_sets() в DataFrame API:

# Spark 3.5+
from pyspark.sql.functions import grouping_sets

result = orders.groupBy("region", "year", "month").agg(
    grouping_sets([("region",), ("year", "month"), ()]),
    F.sum("amount").alias("revenue"),
)

Производительность агрегации

Shuffle Partitions: ключевой параметр

После groupBy Spark создаёт по умолчанию 200 партиций (параметр spark.sql.shuffle.partitions). Это значение оптимально для кластеров уровня Databricks, но не универсально.

# Рекомендация: целевой размер партиции после shuffle - 50–200 MB
# Формула: партиции = (объём данных после shuffle) / 100 MB

# Для небольшой агрегации (результат ~ 1 GB)
spark.conf.set("spark.sql.shuffle.partitions", "20")

# Для большой агрегации (результат ~ 100 GB)
spark.conf.set("spark.sql.shuffle.partitions", "2000")

# AQE автоматически склеивает мелкие партиции (включён по умолчанию в Spark 3.0+)
spark.conf.set("spark.sql.adaptive.enabled", "true")
spark.conf.set("spark.sql.adaptive.coalescePartitions.enabled", "true")

Data Skew: перекос данных в агрегации

Data Skew - ситуация, когда один или несколько ключей группировки содержат несравнимо больше строк, чем остальные. Например, region=MSK с 50 миллионами строк против region=EKB с 200 тысячами.

Диагностика: откройте Spark UI → Stage → Tasks. Если время выполнения задач сильно различается (один task работает в 10–100 раз дольше других) - скорее всего skew.

Решения:

# 1. AQE Skew Join Hint (Spark 3.0+): автоматическое обнаружение и обработка
spark.conf.set("spark.sql.adaptive.skewJoin.enabled", "true")

# 2. Ручное solting: разбить горячий ключ на подгруппы
import random

# Добавить случайный суффикс к горячему ключу → разбить 1 задачу на N
SALT_BUCKETS = 50

salted = orders.withColumn(
    "region_salted",
    F.when(
        col("region") == "MSK",
        F.concat(col("region"), F.lit("_"), (F.rand() * SALT_BUCKETS).cast("int").cast("string"))
    ).otherwise(col("region"))
)

# Агрегация с соленым ключом
partial = salted.groupBy("region_salted").agg(F.sum("amount").alias("partial_revenue"))

# Финальная агрегация: убрать суффикс и сложить части
final = partial.withColumn(
    "region",
    F.regexp_replace(col("region_salted"), r"_\d+$", "")
).groupBy("region").agg(F.sum("partial_revenue").alias("total_revenue"))

Adaptive Query Execution (AQE)

AQE - механизм Spark 3.0+, который динамически адаптирует план выполнения на основе статистики, собранной во время предыдущих стадий.

Три ключевые оптимизации AQE для агрегации:

  1. Coalesce Partitions: после shuffle AQE видит реальные размеры партиций и автоматически склеивает мелкие. Если spark.sql.shuffle.partitions=200, но реальный результат агрегации - 100 MB, AQE объединит 200 партиций в 1–2.

  2. Skew Join: обнаруживает партиции-аутлайеры и автоматически разбивает их на более мелкие подзадачи.

  3. Dynamic Partition Pruning: при Join с фильтрацией одной стороны - передаёт фильтр на scan другой стороны.

# AQE включён по умолчанию в Spark 3.2+
# Проверка
spark.conf.get("spark.sql.adaptive.enabled")  # "true"

# Целевой размер партиции после coalesce
spark.conf.set("spark.sql.adaptive.advisoryPartitionSizeInBytes", "128m")

Explain план агрегации

orders.groupBy("region", "category").agg(
    F.sum("amount").alias("revenue"),
    F.countDistinct("customer_id").alias("customers"),
).explain("extended")

В Physical Plan ищите:

== Physical Plan ==
AdaptiveSparkPlan isFinalPlan=false
+- HashAggregate(keys=[region, category], functions=[sum(amount), count(distinct customer_id)])
   +- Exchange hashpartitioning(region, category, 200), ENSURE_REQUIREMENTS, [plan_id=...]
      +- HashAggregate(keys=[region, category], functions=[partial_sum(amount), partial_count(distinct customer_id)])
         +- FileScan parquet [region, category, amount, customer_id]

Что здесь видно:

  • Два HashAggregate - это и есть двухфазная агрегация (partial + final)
  • Exchange hashpartitioning - shuffle между ними
  • FileScan - Parquet-ридер (если здесь есть PushedFilters - predicate pushdown работает)
  • countDistinct требует дополнительного shuffle по ключу customer_id

Практика: построение Sales Mart

Разберём полный pipeline построения аналитической витрины продаж:

import pyspark.sql.functions as F
from pyspark.sql.functions import col, year, month, quarter

# Чтение данных
orders    = spark.read.parquet("/data/silver/orders/")
products  = spark.read.parquet("/data/silver/products/")
customers = spark.read.parquet("/data/silver/customers/")

# Обогащение данных (join до агрегации - уменьшает shuffle)
enriched = (
    orders
    .filter(col("status") == "COMPLETED")           # фильтр ДО join
    .filter(col("created_at") >= "2024-01-01")      # partitionFilter
    .join(products.select("product_id", "category", "subcategory"), "product_id", "left")
    .join(customers.select("customer_id", "region", "city", "segment"), "customer_id", "left")
    .withColumn("year",    year(col("created_at")))
    .withColumn("month",   month(col("created_at")))
    .withColumn("quarter", quarter(col("created_at")))
)

# --- Агрегат 1: ежемесячная витрина по категориям ---
monthly_by_category = (
    enriched
    .groupBy("year", "month", "region", "category")
    .agg(
        F.sum("amount").alias("revenue"),
        F.count("*").alias("orders"),
        F.countDistinct("customer_id").alias("unique_customers"),
        F.avg("amount").alias("avg_order_value"),
        F.sum(F.when(col("is_new_customer"), col("amount"))).alias("new_customer_revenue"),
    )
)

# --- Агрегат 2: иерархический отчёт с rollup ---
hierarchical = (
    enriched
    .rollup("year", "quarter", "month", "region", "city")
    .agg(
        F.sum("amount").alias("revenue"),
        F.count("*").alias("orders"),
        F.grouping_id(
            "year", "quarter", "month", "region", "city"
        ).alias("level"),
    )
    .withColumn(
        "granularity",
        F.when(col("level") == 0,  "city")
         .when(col("level") == 1,  "region")
         .when(col("level") == 3,  "month")
         .when(col("level") == 7,  "quarter")
         .when(col("level") == 15, "year")
         .otherwise("total")
    )
)

# --- Запись результатов ---
monthly_by_category.write \
    .mode("overwrite") \
    .partitionBy("year", "month") \
    .parquet("/data/gold/sales_mart/monthly_category/")

hierarchical.write \
    .mode("overwrite") \
    .partitionBy("year") \
    .parquet("/data/gold/sales_mart/hierarchical/")

Anti-patterns в агрегациях

Best practices production агрегации

Правило 1: Фильтруй до агрегации

# Плохо: агрегация всех данных, потом фильтр результата
all_regions = orders.groupBy("region").agg(F.sum("amount"))
msk_only    = all_regions.filter(col("region") == "MSK")

# Хорошо: фильтр до groupBy - меньше данных в shuffle
msk_only = orders.filter(col("region") == "MSK") \
    .groupBy("region").agg(F.sum("amount"))

Правило 2: Pre-aggregation при join перед groupBy

# Если нужно объединить два источника и посчитать агрегат:
# Сначала агрегируй каждый источник, потом join агрегатов

# Вместо join (100M rows) → groupBy
# Делай groupBy (100M → 1k) → join (1k) → groupBy (1k)
orders_agg   = orders.groupBy("customer_id").agg(F.sum("amount").alias("spent"))
customer_seg = customers.select("customer_id", "segment")
result       = orders_agg.join(customer_seg, "customer_id")

Правило 3: Контролируй shuffle.partitions

# Для агрегации с маленьким результатом
spark.conf.set("spark.sql.shuffle.partitions", "20")

# После агрегации - всегда проверяй число партиций
result = orders.groupBy("region").agg(...)
result.rdd.getNumPartitions()  # если << 200, AQE уже сработал

Правило 4: Проверяй plan перед деплоем

# Убедиться в: column pruning, predicate pushdown,
# отсутствии лишних exchange, правильном join типе
orders.groupBy("region").agg(F.sum("amount")).explain("extended")

Задания для практики

Задание 1: Многомерная агрегация

Дан DataFrame транзакций с колонками date, region, category, amount, customer_id.

Напишите один agg() запрос по groupBy("region", "category"), который считает:

  • общую выручку
  • количество заказов
  • среднюю сумму заказа
  • медианную сумму (через percentile_approx)
  • количество уникальных клиентов (точное и приближённое)

Сравните время выполнения точного и приближённого countDistinct.

Решение

from pyspark.sql.functions import (
    sum as fsum, count, avg, percentile_approx,
    countDistinct, approx_count_distinct, col
)

transactions = spark.createDataFrame([
    ("2024-01-10", "MSK", "Electronics", 1200.0, 101),
    ("2024-01-11", "MSK", "Electronics",  800.0, 102),
    ("2024-01-12", "MSK", "Clothing",     300.0, 101),
    ("2024-01-13", "SPB", "Electronics", 1500.0, 103),
    ("2024-01-14", "SPB", "Clothing",     400.0, 104),
    ("2024-01-15", "SPB", "Clothing",     250.0, 103),
    ("2024-01-16", "MSK", "Electronics",  950.0, 101),
    ("2024-01-17", "MSK", "Clothing",     120.0, 105),
], ["date", "region", "category", "amount", "customer_id"])

result = (
    transactions
    .groupBy("region", "category")
    .agg(
        fsum("amount").alias("revenue"),
        count("*").alias("orders"),
        avg("amount").alias("avg_order"),
        percentile_approx("amount", 0.5).alias("median_amount"),
        countDistinct("customer_id").alias("unique_customers_exact"),
        approx_count_distinct("customer_id").alias("unique_customers_approx"),
    )
)
result.show()
# +------+-----------+-------+------+---------+-------------+---------------------+-----------------------+
# |region|   category|revenue|orders|avg_order|median_amount|unique_customers_exact|unique_customers_approx|
# +------+-----------+-------+------+---------+-------------+----------------------+-----------------------+
# |   MSK|Electronics| 2950.0|     3|   983.33|        950.0|                     2|                      2|
# |   MSK|   Clothing|  420.0|     2|    210.0|        120.0|                     2|                      2|
# |   SPB|Electronics| 1500.0|     1|   1500.0|       1500.0|                     1|                      1|
# |   SPB|   Clothing|  650.0|     2|    325.0|        250.0|                     1|                      1|
# +------+-----------+-------+------+---------+-------------+----------------------+-----------------------+
Разница в производительности: countDistinct требует дополнительного shuffle для дедупликации - каждое уникальное значение customer_id должно оказаться на одном executor'е, чтобы точно посчитать уникальных. approx_count_distinct использует алгоритм HyperLogLog: каждый executor считает локальный sketch, финальный executor объединяет sketches. Никакого дополнительного shuffle. На 100 млн строк с 10 млн уникальных значений approx_count_distinct в 5-10 раз быстрее при погрешности ≤ 5%.

Задание 2: Rollup для отчёта

Постройте иерархический отчёт по year → quarter → month → region через rollup. Добавьте колонку granularity через grouping_id, которая читаемо описывает уровень агрегации. Отфильтруйте только строки с уровнем детализации до месяца (без детализации по городу).

Решение

from pyspark.sql.functions import (
    sum as fsum, count, grouping_id, when, col,
    year, month, quarter, to_date
)

transactions = spark.createDataFrame([
    ("2024-01-15", "MSK", 1200.0),
    ("2024-01-22", "SPB",  800.0),
    ("2024-02-10", "MSK", 1500.0),
    ("2024-02-14", "SPB",  300.0),
    ("2024-04-05", "MSK",  900.0),
    ("2024-04-18", "SPB",  650.0),
], ["date_str", "region", "amount"]).select(
    year(to_date(col("date_str"), "yyyy-MM-dd")).alias("year"),
    quarter(to_date(col("date_str"), "yyyy-MM-dd")).alias("quarter"),
    month(to_date(col("date_str"), "yyyy-MM-dd")).alias("month"),
    col("region"),
    col("amount"),
)

# rollup("year","quarter","month","region") даёт grouping_id:
#   0 = все 4 колонки (year+quarter+month+region)
#   1 = year+quarter+month (region=null)
#   3 = year+quarter (month=null, region=null)
#   7 = year only
#  15 = grand total (все null)
result = (
    transactions
    .rollup("year", "quarter", "month", "region")
    .agg(
        fsum("amount").alias("revenue"),
        count("*").alias("orders"),
        grouping_id("year", "quarter", "month", "region").alias("gid"),
    )
    .withColumn(
        "granularity",
        when(col("gid") == 0,  "month+region")
        .when(col("gid") == 1,  "month")
        .when(col("gid") == 3,  "quarter")
        .when(col("gid") == 7,  "year")
        .otherwise("grand_total")
    )
    # Оставляем только уровни до месяца: month+region (0) и month (1)
    .filter(col("gid") <= 1)
    .orderBy("year", "quarter", "month", "region")
)
result.show()
# +----+-------+-----+------+-------+------+------------+
# |year|quarter|month|region|revenue|orders| granularity|
# +----+-------+-----+------+-------+------+------------+
# |2024|      1|    1|   MSK| 1200.0|     1|month+region|
# |2024|      1|    1|   SPB|  800.0|     1|month+region|
# |2024|      1|    1|  null| 2000.0|     2|       month|
# |2024|      1|    2|   MSK| 1500.0|     1|month+region|
# |2024|      1|    2|   SPB|  300.0|     1|month+region|
# |2024|      1|    2|  null| 1800.0|     2|       month|
# |2024|      2|    4|   MSK|  900.0|     1|month+region|
# |2024|      2|    4|   SPB|  650.0|     1|month+region|
# |2024|      2|    4|  null| 1550.0|     2|       month|
# +----+-------+-----+------+-------+------+------------+
grouping_id кодирует отсутствующие ключи как биты: бит 0 = последняя колонка (region), бит 1 = month, и т.д. При gid=1 только region - null, то есть агрегация на уровне месяца без разбивки по региону. Фильтр gid <= 1 оставляет только детализацию месяц+регион и просто месяц.

Задание 3: Диагностика skew

Создайте синтетический DataFrame с сильным перекосом: один ключ имеет 90% всех строк. Запустите groupBy().agg() и через Spark UI найдите задачу-аутлайер. Примените salting и убедитесь, что время выполнения сократилось.

Решение

from pyspark.sql.functions import (
    col, sum as fsum, concat, lit, rand, count
)

# Синтетический датасет с сильным перекосом
# 100 ключей × 100 строк = 10 000 строк (нормальные)
# 1 горячий ключ × 90 000 строк (скошенные)
normal = spark.range(10_000).select(
    concat(lit("key_"), (col("id") % 100).cast("string")).alias("key"),
    (col("id") * 7 % 1000).cast("double").alias("amount"),
)
hot = spark.range(90_000).select(
    lit("key_HOT").alias("key"),
    (col("id") * 3 % 1000).cast("double").alias("amount"),
)
df = normal.union(hot).repartition(10)  # 10 исходных партиций

# --- Без salting: задача-аутлайер обрабатывает 90% данных ---
result_naive = (
    df.groupBy("key")
      .agg(fsum("amount").alias("total"), count("*").alias("cnt"))
)
# result_naive.count()  # → в Spark UI одна задача займёт >>остальных

# --- С salting: размазываем hot key по N партициям ---
N_SALT = 10
salted = df.withColumn(
    "salted_key",
    concat(col("key"), lit("_"), (rand() * N_SALT).cast("int").cast("string"))
)

# Первичная агрегация: по salted_key (хранит оригинальный key как поле)
partial = (
    salted
    .groupBy("salted_key", "key")
    .agg(
        fsum("amount").alias("partial_total"),
        count("*").alias("partial_cnt"),
    )
)

# Финальная агрегация: восстанавливаем оригинальный ключ
result_salted = (
    partial
    .groupBy("key")
    .agg(
        fsum("partial_total").alias("total"),
        fsum("partial_cnt").alias("cnt"),
    )
)
result_salted.orderBy(col("cnt").desc()).show(5)
# +-------+----------+-----+
# |    key|     total|  cnt|
# +-------+----------+-----+
# |key_HOT|...|90000|   ← был аутлайер, теперь split на 10
# |  key_0|...|  100|
# +-------+----------+-----+
Механизм: без salting все 90 000 строк key_HOT попадают на один executor после shuffle (partition key hash), остальные executor'ы заканчивают за секунды и ждут. С salting добавляем случайный суффикс _0.._9 - горячий ключ разбивается на 10 ключей, каждый executor обрабатывает ~9 000 строк. Цена: двухфазная агрегация (2 shuffle вместо 1). На практике включается автоматически через AQE Skew Join Optimization (spark.sql.adaptive.skewJoin.enabled=true), но для groupBy salting нужно применять вручную.

Задание 4: Explain план

Напишите агрегацию с countDistinct и без. Сравните их explain("extended"). Найдите различие в физическом плане: почему countDistinct требует дополнительного Exchange.

Решение

from pyspark.sql.functions import count, countDistinct, approx_count_distinct, col, sum as fsum

df = spark.createDataFrame(
    [(i % 20, i % 500, float(i)) for i in range(5_000)],
    ["key", "user_id", "amount"]
)

# --- Вариант A: простая агрегация (count + sum) ---
simple = df.groupBy("key").agg(
    count("*").alias("orders"),
    fsum("amount").alias("revenue"),
)
print("=== Простая агрегация ===")
simple.explain()
# == Physical Plan ==
# *(2) HashAggregate(keys=[key#0], functions=[count(1), sum(amount#2)])
# +- Exchange hashpartitioning(key#0, 200), ENSURE_REQUIREMENTS
#    +- *(1) HashAggregate(keys=[key#0], functions=[partial_count(1), partial_sum(amount#2)])
#       +- *(1) LocalTableScan [key#0, user_id#1, amount#2]
# Два HashAggregate (partial + final) и один Exchange - стандартная двухфазная схема.

# --- Вариант B: с countDistinct ---
with_distinct = df.groupBy("key").agg(
    count("*").alias("orders"),
    countDistinct("user_id").alias("unique_users"),
)
print("=== С countDistinct ===")
with_distinct.explain()
# == Physical Plan ==
# *(3) HashAggregate(keys=[key#0], functions=[count(1), count(user_id#1)])
# +- Exchange hashpartitioning(key#0, 200), ENSURE_REQUIREMENTS
#    +- *(2) HashAggregate(keys=[key#0], functions=[partial_count(1), partial_count(user_id#1)])
#       +- *(2) Filter isnotnull(user_id#1)
#          +- *(1) LocalTableScan [key#0, user_id#1]
# ВАЖНО: partial_count(user_id#1) здесь означает дедупликацию внутри каждой
# partition ДО shuffle. После shuffle каждый executor видит все значения
# user_id для своего key и считает точное количество уникальных.

# --- Вариант C: approx_count_distinct (без дополнительного shuffle) ---
with_approx = df.groupBy("key").agg(
    count("*").alias("orders"),
    approx_count_distinct("user_id").alias("approx_unique_users"),
)
print("=== С approx_count_distinct ===")
with_approx.explain()
# В плане: partial_approx_count_distinct → Exchange → merge sketches
# Каждый executor строит HyperLogLog-sketch локально (маленький объект ~1 KB),
# финальный executor объединяет sketches. Никакой передачи сырых user_id по сети.
Ключевое различие в планах:

Метод Частичная агрегация Shuffle трафик Финальная агрегация
count(*) partial_count (счётчик) только ключи + счётчики суммирует счётчики
countDistinct частичная дедупликация ключи + уникальные значения per partition точный подсчёт
approx_count_distinct HyperLogLog sketch ключи + sketch (~1 KB) объединение sketches

countDistinct не может полностью свернуть данные на mapper - он должен передать все уникальные значения через shuffle, чтобы executor с финальной агрегацией мог точно посчитать уникальных. Чем выше cardinality user_id - тем больше shuffle трафик и тем сильнее разница с approx_count_distinct.