Data Skew: ручной salting - генерация ключей и двухпроходная агрегация

Ручные техники борьбы с Data Skew когда AQE бессилен: точечная диагностика hot keys, Random и Deterministic Salting, двухпроходная агрегация для groupBy, salting-JOIN с explode для справочников, изолированное сальтирование топ-ключей, анализ в Spark UI и полная матрица выбора стратегии.

optimization

1. Границы применимости AQE и природа неуправляемого Data Skew

Предыдущий урок разобрал как AQE Skew Join автоматически разбивает скошенные Shuffle-партиции на части. Это мощный инструмент, но у него есть принципиальные ограничения. Данный урок посвящён ситуациям, когда автоматика бессильна и инженер обязан вмешаться вручную.

Где пасует AQE: три сценария

Сценарий 1: groupBy с доминирующим ключом. AQE Skew Join работает только для JOIN-операций. Операция groupBy("country").sum("revenue") не является JOIN - это агрегация. AQE не может «разбить» горячий ключ GROUP BY на части, потому что агрегация по своей природе требует, чтобы все строки с одинаковым ключом попали в одну функцию-агрегатор.

Если 85% строк имеют country = "RU", все они должны встретиться на одном Executor'е для вычисления суммы. AQE бессилен.

Сценарий 2: Full Outer Join. AQE Skew Join не поддерживает FULL OUTER JOIN из-за сложности генерации NULL-строк при отсутствии совпадений. Если такой JOIN скошен - AQE молча пропускает оптимизацию.

Сценарий 3: Экстремальный skew, превышающий физическую память. AQE разбивает скошенную партицию на N частей по advisoryPartitionSizeInBytes (например, 64 MB). Но если одна строка содержит гигантский вложенный массив (ARRAY with 10 million elements), или если после разбиения каждый Split всё равно не умещается в Executor Memory - AQE не поможет. Container упадёт с OOM.

Анатомия критического OOM при Data Skew

Схема показывает физику OOM: суммарный объём данных для одного ключа, попавший на один Executor, может в десятки раз превысить выделенную ему память. Spill на диск замедляет выполнение в сотни раз, а если даже spill не помогает - YARN/K8s убивает контейнер.

Симптомы в логах: как распознать OOM от skew

# Симптом 1: YARN убивает контейнер из-за превышения памяти
ERROR YarnAllocator: Container killed by YARN for exceeding limits.
  11.2 GB of 8 GB physical memory used.
  Consider boosting spark.executor.memoryOverhead.

# Симптом 2: Java GC не справляется
ERROR Executor: Exception in task 199.0 in stage 3.0 (TID 1299)
java.lang.OutOfMemoryError: GC overhead limit exceeded
    at org.apache.spark.unsafe.types.UTF8String...

# Симптом 3: Spill на диск (медленно, но ещё работает)
INFO ExternalSorter: Task 199: using disk storage for spill
  spill_size=12.3 GB

# Симптом 4: В Spark UI → Stage → Tasks
# Task 199: Duration=2.5h, GC Time=45min, Spill(memory)=45GB

Если видите какой-либо из этих симптомов - вы имеете дело с Data Skew, требующим ручного вмешательства.


2. Диагностика: находим горячие ключи до оптимизации

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

Метод 1: быстрый подсчёт частот

from pyspark.sql import SparkSession, functions as F

spark = SparkSession.builder.getOrCreate()

# Загружаем данные (без action - plan не выполняется)
df = spark.table("bronze.user_events")

# Считаем частоту каждого значения ключа группировки
# approxCountDistinct - быстрее точного count distinct
key_distribution = df.groupBy("user_id") \
    .count() \
    .orderBy(F.desc("count"))

# Топ-20 горячих ключей
print("Топ-20 горячих ключей:")
key_distribution.limit(20).show()

# Статистика распределения
stats = key_distribution.agg(
    F.count("count").alias("unique_keys"),
    F.sum("count").alias("total_rows"),
    F.max("count").alias("max_rows_per_key"),
    F.min("count").alias("min_rows_per_key"),
    F.percentile_approx("count", 0.5).alias("median_rows_per_key"),
    F.percentile_approx("count", 0.99).alias("p99_rows_per_key"),
)
stats.show()

# Ключевые метрики для решения:
# - skew_ratio = max / median (если > 5: AQE может помочь, если > 100: нужен salting)
# - top_key_pct = max / total_rows * 100 (если > 50%: критический skew)

Метод 2: approxQuantile для больших таблиц

# approxQuantile работает значительно быстрее полного groupBy на больших таблицах
# Но требует числового типа ключа

# Для числовых ключей:
quantiles = df.stat.approxQuantile(
    "user_id",
    [0.25, 0.5, 0.75, 0.99, 0.999],
    0.01  # относительная погрешность
)
print(f"Quantile 25%:   {quantiles[0]}")
print(f"Quantile 50%:   {quantiles[1]}")
print(f"Quantile 99%:   {quantiles[3]}")
print(f"Quantile 99.9%: {quantiles[4]}")

# Для строковых ключей: используем freqItems (топ частых значений)
frequent = df.stat.freqItems(["user_id"], support=0.01)
# support=0.01: вернёт ключи встречающиеся в >= 1% строк
frequent.show(truncate=False)

Метод 3: диагностика через Spark UI без полного scan

# Используем sample для быстрой оценки без полного scan больших таблиц
# Точность: достаточна для выявления горячих ключей с > 1% данных

sample_fraction = 0.01  # 1% выборка

sample_dist = df.sample(fraction=sample_fraction, seed=42) \
    .groupBy("user_id") \
    .count() \
    .orderBy(F.desc("count")) \
    .limit(50)

print(f"Распределение по 1% выборке (scale ×{int(1/sample_fraction)}x):")
sample_dist.withColumn(
    "estimated_total",
    F.col("count") * int(1 / sample_fraction)
).show()

Интерпретация результатов: когда какой инструмент применять

def diagnose_and_recommend(max_count: int, median_count: int, total_rows: int) -> str:
    """
    На основе статистики распределения ключей рекомендует стратегию.
    """
    skew_ratio = max_count / max(median_count, 1)
    top_key_pct = max_count / total_rows * 100

    if skew_ratio < 5:
        return "Skew незначительный. AQE с дефолтными настройками достаточен."
    elif skew_ratio < 50 and top_key_pct < 20:
        return "Умеренный skew. Рекомендуется AQE с уменьшенным skewedPartitionFactor=3."
    elif skew_ratio < 200 and top_key_pct < 50:
        return "Значительный skew. AQE Skew Join + точечный salting для топ-3 ключей."
    elif top_key_pct > 50:
        return "КРИТИЧЕСКИЙ SKEW! Топ ключ содержит >50% данных. Обязателен ручной salting."
    else:
        return "Тяжёлый skew. Ручной salting + двухпроходная агрегация."

3. Концепция Salting: математика паттерна

Salting (сальтирование) - это техника искусственного увеличения кардинальности ключа. Вместо одного значения key=X мы создаём N значений key=X_0, key=X_1, ..., key=X_(N-1). Spark хэш-функция теперь распределяет эти N значений по N разным партициям.

Почему это работает: математика хэш-распределения

Без salting:
  hash("customer_123") % 200 = 47  ← всегда одна партиция!
  50M строк → Partition 47 → один Executor → OOM

С salting (N=100):
  hash("customer_123_0") % 200 = 47   ← 500K строк
  hash("customer_123_1") % 200 = 183  ← 500K строк
  hash("customer_123_2") % 200 = 12   ← 500K строк
  ...
  hash("customer_123_99") % 200 = 91  ← 500K строк
  50M строк → 100 разных партиций → 100 Executor'ов работают параллельно

Как правильно выбрать Salting Factor N

Salting Factor - это число «копий» каждого горячего ключа. Выбор N имеет значение:

  • N слишком маленький (например, N=5): горячий ключ разбивается только на 5 партиций. Если одна партиция содержит 50 GB данных → каждая из 5 частей = 10 GB. При Executor Memory 8 GB всё равно OOM.
  • N оптимальный: ceil(hot_partition_size / executor_memory). Если партиция 50 GB и Executor 8 GB → N = ceil(50/8) = 7. С запасом берём N=10.
  • N слишком большой (например, N=10000): справочник right-table расширяется в 10000 раз → гигантский overhead памяти и сети.
def calculate_salt_factor(
    hot_partition_estimated_gb: float,
    executor_memory_gb: float,
    safety_margin: float = 0.7,  # используем 70% памяти Executor
) -> int:
    """
    Вычисляет минимально необходимый Salting Factor.

    safety_margin: доля памяти Executor для обработки данных.
    Остальная память нужна для JVM overhead, GC, shuffle buffers.
    """
    usable_memory = executor_memory_gb * safety_margin
    factor = hot_partition_estimated_gb / usable_memory
    # Округляем вверх и берём ближайшую степень 10 для чистоты
    recommended = max(2, int(factor) + 1)
    return recommended

# Примеры:
print(calculate_salt_factor(50, 8))   # 50GB горячая партиция, 8GB Executor → 9
print(calculate_salt_factor(200, 16)) # 200GB, 16GB → 18
print(calculate_salt_factor(5, 8))    # 5GB, 8GB → 2 (минимальный)

4. Random Salting для JOIN: полная реализация

Самый частый сценарий использования salting - это JOIN скошенной fact-таблицы с dimension-таблицей (справочником).

Структура задачи

fact_table (orders): 500M строк
  customer_id: 90% строк = customer_0 (гостевые заказы)

dimension_table (customers): 10K строк
  customer_id, customer_name, tier

Без salting: JOIN по customer_id гарантирует что все 450M строк с customer_0 попадут на один Executor.

Ключевое ограничение: почему нельзя просто «посолить» обе стороны

Соль - это случайное число. Если добавить случайную соль к обеим таблицам независимо, строки с одинаковым customer_id могут получить разные значения соли и не совпасть при JOIN. Результат JOIN будет неправильным.

Правило: соль добавляется только к одной (скошенной) стороне. Другая сторона расширяется (explode) - для каждой строки справочника создаётся N копий, по одной для каждого значения соли.

Реализация шаг за шагом

from pyspark.sql import SparkSession, functions as F
from pyspark.sql.types import IntegerType, ArrayType

spark = SparkSession.builder \
    .config("spark.sql.adaptive.enabled", "true")  \
    .getOrCreate()

# Исходные данные
orders = spark.table("bronze.orders")
customers = spark.table("silver.customers")

# Параметры salting
SALT_FACTOR = 50  # 50 копий горячего ключа

# ════════════════════════════════════════════════════════════════
# ШАГ 1: Добавляем случайную соль к fact-таблице (orders)
# ════════════════════════════════════════════════════════════════
# rand() генерирует равномерное float [0, 1)
# Умножаем на SALT_FACTOR и конвертируем в int → [0, SALT_FACTOR-1]
# Это наша "соль" - случайное число для каждой строки

orders_salted = orders.withColumn(
    "salt",
    (F.rand() * SALT_FACTOR).cast(IntegerType())
).withColumn(
    # Составной ключ: customer_id + "_" + salt
    # "customer_0_23", "customer_0_7", "customer_0_41", ...
    "customer_id_salted",
    F.concat_ws("_", F.col("customer_id"), F.col("salt"))
)

# ════════════════════════════════════════════════════════════════
# ШАГ 2: Расширяем dimension-таблицу (customers)
# ════════════════════════════════════════════════════════════════
# Каждую строку customers нужно превратить в SALT_FACTOR копий
# каждая с своим суффиксом соли

# Создаём массив [0, 1, 2, ..., SALT_FACTOR-1]
salt_array = F.array([F.lit(i) for i in range(SALT_FACTOR)])
# Примечание: для больших SALT_FACTOR используйте:
# spark.range(SALT_FACTOR).select(F.col("id").cast("int").alias("salt"))
# и crossJoin - это эффективнее чем большой array literal

customers_expanded = customers.withColumn(
    # explode превращает массив [0,1,...,49] в отдельные строки
    # Каждая строка customers → SALT_FACTOR строк в customers_expanded
    "salt",
    F.explode(salt_array)
).withColumn(
    "customer_id_salted",
    F.concat_ws("_", F.col("customer_id"), F.col("salt"))
)

# ════════════════════════════════════════════════════════════════
# ШАГ 3: JOIN по составному ключу (customer_id_salted)
# ════════════════════════════════════════════════════════════════
# Теперь Spark видит 50 разных значений вместо одного hot key
# Данные равномерно распределяются по партициям

result = orders_salted.join(
    customers_expanded,
    on="customer_id_salted",
    how="left"
)

# ════════════════════════════════════════════════════════════════
# ШАГ 4: Очистка технических колонок
# ════════════════════════════════════════════════════════════════
result_clean = result.drop("salt", "customer_id_salted")

# Финальная агрегация (теперь данные уже равномерны!)
final = result_clean.groupBy("customer_id", "customer_name", "tier") \
    .agg(
        F.sum("amount").alias("total_amount"),
        F.count("*").alias("order_count")
    )

final.show(20)

Оптимизированная версия с crossJoin для большого SALT_FACTOR

# Эффективнее для SALT_FACTOR > 20: не создаём огромный array literal
def expand_dimension_table(dim_df, salt_factor: int):
    """
    Расширяет dimension-таблицу в salt_factor раз через crossJoin.
    Более эффективно чем explode(array(0,1,...,N)) для больших N.

    Сложность: O(dim_size × salt_factor) строк в памяти и на диске.
    Для customers с 10K строк и salt_factor=100: 1M строк → допустимо.
    Для customers с 1M строк и salt_factor=100: 100M строк → слишком много!
    В последнем случае используйте Broadcast + Random Salting без explode.
    """
    salt_values = spark.range(salt_factor).select(
        F.col("id").cast(IntegerType()).alias("salt")
    )

    return dim_df.crossJoin(salt_values) \
        .withColumn(
            "join_key",
            F.concat_ws("_", F.col("customer_id"), F.col("salt"))
        )

# Для fact таблицы:
orders_with_key = orders.withColumn(
    "salt", (F.rand() * SALT_FACTOR).cast(IntegerType())
).withColumn(
    "join_key",
    F.concat_ws("_", F.col("customer_id"), F.col("salt"))
)

customers_expanded = expand_dimension_table(customers, SALT_FACTOR)

result = orders_with_key.join(customers_expanded, "join_key", "left") \
    .drop("salt", "join_key")

5. Двухпроходная агрегация: борьба со Skew в groupBy

JOIN - не единственный источник skew. Операции groupBy с агрегацией могут страдать не меньше. Именно для них AQE бессилен - и именно для них Two-Phase Aggregation является элегантным решением.

Проблема скошенного groupBy

Представьте: вам нужно посчитать выручку по каждому пользователю:

# Наивная реализация (вызывает OOM при skew)
revenue_by_user = orders.groupBy("user_id").agg(
    F.sum("amount").alias("total_revenue"),
    F.count("*").alias("order_count")
)

Если user_id = 0 (гость) встречается 90% строк, Spark отправит все эти строки на один Executor для агрегации. Executor получает 90% данных всей таблицы → OOM.

Идея двухпроходной агрегации

Ключевое озарение: после Фазы 1 вместо 100M строк у нас 50 строк промежуточных агрегатов. Фаза 2 работает с 50 строками - это мгновенно, без OOM, без skew.

Полная реализация для разных агрегатных функций

from pyspark.sql import SparkSession, functions as F
from pyspark.sql.types import IntegerType

spark = SparkSession.builder.getOrCreate()

orders = spark.table("bronze.orders")
SALT_FACTOR = 50

# ════════════════════════════════════════════════════════════════
# СЛУЧАЙ 1: SUM - простая двухпроходная агрегация
# ════════════════════════════════════════════════════════════════

# Фаза 1: Добавляем соль + частичная агрегация
orders_salted = orders.withColumn(
    "salt", (F.rand() * SALT_FACTOR).cast(IntegerType())
)

# Частичный SUM: суммируем внутри каждой (user_id, salt) группы
phase1_sum = orders_salted \
    .groupBy("user_id", "salt") \
    .agg(
        F.sum("amount").alias("partial_sum"),
        F.count("*").alias("partial_count")
    )
# Результат: N_unique_user_ids × SALT_FACTOR строк
# Данные равномерно распределены (нет скоса)

# Фаза 2: Убираем соль + финальная агрегация
# Работаем с маленьким промежуточным датасетом
final_sum = phase1_sum \
    .groupBy("user_id") \
    .agg(
        F.sum("partial_sum").alias("total_amount"),   # SUM of SUMs = полный SUM
        F.sum("partial_count").alias("total_count")   # SUM of COUNTs = полный COUNT
    )


# ════════════════════════════════════════════════════════════════
# СЛУЧАЙ 2: COUNT DISTINCT - требует особого подхода
# ════════════════════════════════════════════════════════════════
# Проблема: COUNT DISTINCT нельзя вычислить через SUM(partial COUNT DISTINCT)
# Решение: использовать approx_count_distinct (HyperLogLog) или collect_set

# Вариант А: approxCountDistinct (точность ~1-5%, очень быстро)
phase1_distinct = orders_salted \
    .groupBy("user_id", "salt") \
    .agg(
        F.approx_count_distinct("session_id", rsd=0.01).alias("partial_distinct_sessions")
    )

# НЕЛЬЗЯ: final = phase1_distinct.groupBy("user_id").sum("partial_distinct")
# SUM(HyperLogLog estimates) ≠ correct COUNT DISTINCT

# Вариант Б: Точный COUNT DISTINCT через collect_set (только для небольших множеств)
phase1_collect = orders_salted \
    .groupBy("user_id", "salt") \
    .agg(
        F.collect_set("session_id").alias("partial_sessions")
    )

final_distinct = phase1_collect \
    .groupBy("user_id") \
    .agg(
        # Объединяем все множества сессий, берём union, считаем уникальные
        F.size(
            F.array_distinct(F.flatten(F.collect_list("partial_sessions")))
        ).alias("unique_sessions")
    )
# Предупреждение: collect_list создаёт массив всех элементов в памяти Driver
# Безопасно только если итоговое число уникальных значений мало (< 1M)


# ════════════════════════════════════════════════════════════════
# СЛУЧАЙ 3: AVG - вычисляем через SUM/COUNT
# ════════════════════════════════════════════════════════════════
# AVG нельзя вычислить как AVG(partial AVGs) без весов
# Правильно: AVG = total_sum / total_count

phase1_avg = orders_salted \
    .groupBy("user_id", "salt") \
    .agg(
        F.sum("amount").alias("partial_sum"),
        F.count("*").alias("partial_count")
    )

final_avg = phase1_avg \
    .groupBy("user_id") \
    .agg(
        (F.sum("partial_sum") / F.sum("partial_count")).alias("avg_amount")
    )


# ════════════════════════════════════════════════════════════════
# СЛУЧАЙ 4: MAX/MIN - коммутативны, двухпроходная агрегация тривиальна
# ════════════════════════════════════════════════════════════════
phase1_max = orders_salted \
    .groupBy("user_id", "salt") \
    .agg(
        F.max("amount").alias("partial_max"),
        F.min("amount").alias("partial_min")
    )

final_max_min = phase1_max \
    .groupBy("user_id") \
    .agg(
        F.max("partial_max").alias("max_amount"),   # MAX of MAXes = полный MAX
        F.min("partial_min").alias("min_amount")    # MIN of MINs = полный MIN
    )

6. Изолированное (точечное) сальтирование: оптимизация паттерна

Тотальное сальтирование - применять соль ко всем строкам таблицы - это антипаттерн. Оно создаёт ненужный overhead для «нормальных» ключей, у которых нет проблемы skew. Изолированное сальтирование применяет соль только к горячим ключам.

Почему тотальный salting вреден

При SALT_FACTOR=50 и тотальном сальтировании:

  • Справочник customers (10K строк) расширяется до 10K × 50 = 500K строк (OK)
  • Но если в fact-таблице 100K уникальных ключей, каждый со средним числом строк: создаётся 100K × 50 = 5M групп в Phase 1 вместо 100K → ненужные overhead Shuffle и памяти

Реализация точечного salting через when/otherwise

# Шаг 1: Определяем горячие ключи
hot_keys_df = orders.groupBy("user_id").count() \
    .filter(F.col("count") > 1_000_000) \
    .select("user_id")

hot_keys = {row["user_id"] for row in hot_keys_df.collect()}
print(f"Найдено горячих ключей: {len(hot_keys)}")
print(f"Горячие ключи: {hot_keys}")
# Пример: {0, None, -1}

# Шаг 2: Применяем соль только к горячим ключам
SALT_FACTOR = 100

def apply_targeted_salt(df, hot_key_set: set, key_col: str, salt_factor: int):
    """
    Применяет соль только к горячим ключам.
    Остальные ключи получают суффикс "_0" (единственная версия).

    Это избегает лишнего расширения справочника для нормальных ключей:
    - Горячий ключ user_id=0 → user_id_0_0, user_id_0_1, ..., user_id_0_99
    - Нормальный ключ user_id=12345 → user_id_12345_0 (только одна копия)
    """
    # Создаём условие: является ли ключ горячим
    is_hot = F.col(key_col).isin(list(hot_key_set)) | F.col(key_col).isNull()

    return df.withColumn(
        "salt",
        F.when(is_hot, (F.rand() * salt_factor).cast(IntegerType()))
         .otherwise(F.lit(0))  # нормальные ключи всегда получают соль=0
    ).withColumn(
        "join_key",
        F.concat_ws("_",
                    F.coalesce(F.col(key_col).cast("string"), F.lit("NULL")),
                    F.col("salt").cast("string"))
    )

orders_salted = apply_targeted_salt(orders, hot_keys, "user_id", SALT_FACTOR)

# Шаг 3: Расширяем справочник только для горячих ключей
def expand_for_hot_keys(dim_df, hot_key_set: set, key_col: str, salt_factor: int):
    """
    Расширяет dimension-таблицу:
    - Горячие ключи → salt_factor копий каждой строки
    - Нормальные ключи → 1 копия с суффиксом _0

    Итого: (normal_rows × 1) + (hot_rows × salt_factor) строк
    Это значительно меньше чем (all_rows × salt_factor) при тотальном salting.
    """
    is_hot = F.col(key_col).isin(list(hot_key_set)) | F.col(key_col).isNull()

    # Горячие строки: расширяем
    hot_rows = dim_df.filter(is_hot)
    salt_array = F.array([F.lit(i) for i in range(salt_factor)])
    hot_expanded = hot_rows.withColumn("salt", F.explode(salt_array))

    # Нормальные строки: только с солью=0
    normal_rows = dim_df.filter(~is_hot).withColumn("salt", F.lit(0))

    expanded = hot_expanded.union(normal_rows) \
        .withColumn(
            "join_key",
            F.concat_ws("_",
                        F.coalesce(F.col(key_col).cast("string"), F.lit("NULL")),
                        F.col("salt").cast("string"))
        )

    return expanded

customers_expanded = expand_for_hot_keys(customers, hot_keys, "user_id", SALT_FACTOR)

# Шаг 4: JOIN по join_key
result = orders_salted.join(customers_expanded, "join_key", "left") \
    .drop("salt", "join_key")

print(f"Нормальные строки customers: {customers.count()}")
print(f"Расширенные строки customers: {customers_expanded.count()}")
print(f"Overhead: {customers_expanded.count() / customers.count():.1f}x")
# Вместо 100x (тотальный) получаем ~3x (только горячие ключи расширяются)

Детерминированное Salting: воспроизводимые результаты

Random Salting имеет ограничение: при повторном запуске результаты могут немного отличаться (если агрегация зависит от порядка). Для ETL-пайплайнов где важна воспроизводимость используется Deterministic Salting:

def deterministic_salt(df, key_col: str, salt_factor: int):
    """
    Deterministic Salting: соль вычисляется из хэша другой колонки.
    Результат воспроизводим между запусками.

    Применяется когда:
    1. Нужна идемпотентность (инкрементальный ETL, CDC)
    2. Результаты должны совпадать при повторном запуске
    3. Salt определяется бизнес-ключом (например, order_id)
    """
    return df.withColumn(
        "salt",
        # pmod(hash(order_id), salt_factor) → детерминированный [0, salt_factor-1]
        # hash() в Spark возвращает MurmurHash - быстро и равномерно
        F.pmod(F.hash(F.col("order_id")), F.lit(salt_factor)).cast(IntegerType())
    ).withColumn(
        "join_key",
        F.concat_ws("_",
                    F.coalesce(F.col(key_col).cast("string"), F.lit("NULL")),
                    F.col("salt").cast("string"))
    )

# Пример: два запуска дадут одинаковые результаты
orders_det1 = deterministic_salt(orders, "user_id", 50)
orders_det2 = deterministic_salt(orders, "user_id", 50)
# orders_det1 и orders_det2 идентичны - salt вычисляется из order_id

7. Анализ планов выполнения и Spark UI для salted запросов

Что изменяется в Physical Plan при salting

# Смотрим план ДО salting (наивный groupBy):
orders.groupBy("user_id").sum("amount").explain("formatted")
== Physical Plan ==
HashAggregate(keys=[user_id#1], functions=[sum(amount#3)])
+- Exchange hashpartitioning(user_id#1, 200), ENSURE_REQUIREMENTS
   +- HashAggregate(keys=[user_id#1], functions=[partial_sum(amount#3)])
      +- Scan parquet bronze.orders

# Видим ONE shuffle для user_id: все user_id=0 → Partition 47 → OOM
# После двухпроходной агрегации:
phase1_sum.explain("formatted")
== Physical Plan ==
HashAggregate(keys=[user_id#1, salt#99], functions=[sum(amount#3), count(1)])
+- Exchange hashpartitioning(user_id#1, salt#99, 200), ENSURE_REQUIREMENTS
   +- HashAggregate(keys=[user_id#1, salt#99], functions=[partial_sum(amount#3)])
      +- Project [user_id#1, amount#3, (rand(seed=...) * 50.0) AS salt#99]
         +- Scan parquet bronze.orders

# Видим: группировка по (user_id, salt) → данные равномерно распределены!
# user_id=0 теперь → 50 разных ключей по разным партициям
# При salted JOIN с explode:
result.explain("formatted")
== Physical Plan ==
Project [...]
+- SortMergeJoin [join_key#100], [join_key#201], LeftOuter
   :- Sort [join_key#100 ASC], false, 0
   :  +- Exchange hashpartitioning(join_key#100, 200)
   :     +- Project [... join_key = concat(user_id, "_", salt) ...]
   :        +- Scan parquet bronze.orders
   +- Sort [join_key#201 ASC], false, 0
      +- Exchange hashpartitioning(join_key#201, 200)
         +- Generate explode(array(0, 1, ..., 49)), ...  ← explode расширяет customers
            +- Scan parquet silver.customers

# Generate узел (explode) виден в плане - это расширение справочника

Что искать в Spark UI для верификации успеха

Ключевые метрики в Spark UI → Stage → Summary Metrics:

  • Max Task Duration резко снизилась: было 4h, стало 6s
  • Task Duration Distribution стала равномерной: все tasks ~одного времени
  • Shuffle Spill (Memory) = 0 (раньше был гигантский)
  • GC Time снизилось до нормальных 2-5% от Duration (раньше 40-60%)

В вкладке JobsStages:

  • Появился дополнительный Stage (Фаза 1 → Фаза 2 в двухпроходной агрегации)
  • В Phase 1 Stage: Tasks распределены равномерно, Max≈Median
  • В Phase 2 Stage: очень быстрый Stage (работает с маленьким промежуточным результатом)

Матрица выбора стратегии борьбы со Data Skew

Подводим итог: какой инструмент применять в какой ситуации.

Полный пример: ETL-пайплайн с комплексной защитой от skew

# complete_skew_resistant_pipeline.py
# Демонстрирует применение нескольких техник в одном пайплайне

from pyspark.sql import SparkSession, functions as F
from pyspark.sql.types import IntegerType
import time

spark = SparkSession.builder \
    .master("local[8]") \
    .appName("skew-resistant-etl") \
    .config("spark.sql.adaptive.enabled", "true") \
    .config("spark.sql.adaptive.skewJoin.enabled", "true") \
    .getOrCreate()

spark.sparkContext.setLogLevel("ERROR")


def create_skewed_pipeline_data(spark: SparkSession):
    """Создаёт реалистичный скошенный датасет для ETL-пайплайна."""
    n_events = 3_000_000

    # Кликстрим: 85% событий от незалогиненных пользователей
    events = spark.range(n_events).select(
        F.col("id").alias("event_id"),
        F.when(F.rand() < 0.85, F.lit(0))
         .when(F.rand() < 0.10, F.lit(1))
         .otherwise((F.rand() * 98 + 2).cast("int"))
         .alias("user_id"),
        F.array(F.lit("view"), F.lit("click"), F.lit("purchase")).getItem(
            (F.rand() * 3).cast("int")
        ).alias("event_type"),
        (F.rand() * 200).alias("value"),
        F.date_add(F.lit("2024-01-01"), (F.rand() * 30).cast("int")).alias("date")
    )

    # Таблица пользователей
    users = spark.range(100).select(
        F.col("id").alias("user_id"),
        F.array(F.lit("free"), F.lit("premium"), F.lit("enterprise")).getItem(
            (F.col("id") % 3).cast("int")
        ).alias("subscription"),
        F.concat(F.lit("User_"), F.col("id")).alias("username")
    )

    return events, users


# ── ЭТАП 1: Диагностика ────────────────────────────────────────────────────
print("=== Диагностика распределения ===")
events, users = create_skewed_pipeline_data(spark)
events.cache()
users.cache()
events.count()
users.count()

top_keys = events.groupBy("user_id") \
    .count() \
    .orderBy(F.desc("count")) \
    .limit(5) \
    .collect()
for row in top_keys:
    print(f"  user_id={row.user_id}: {row['count']:,} событий")

hot_keys = {row.user_id for row in top_keys if row['count'] > 100_000}
SALT = 20

# ── ЭТАП 2: Двухпроходная агрегация для groupBy ────────────────────────────
print("\n=== Этап 2: Метрики событий по пользователю (Two-Phase) ===")
t0 = time.time()

# Фаза 1: Частичная агрегация с солью
events_salted = events.withColumn(
    "salt", (F.rand() * SALT).cast(IntegerType())
)

phase1 = events_salted \
    .groupBy("user_id", "salt", "event_type") \
    .agg(
        F.sum("value").alias("partial_value"),
        F.count("*").alias("partial_count")
    )

# Фаза 2: Финальная агрегация
user_metrics = phase1 \
    .groupBy("user_id", "event_type") \
    .agg(
        F.sum("partial_value").alias("total_value"),
        F.sum("partial_count").alias("total_events")
    )

metrics_count = user_metrics.count()
t_phase = time.time() - t0
print(f"  Two-Phase Aggregation: {t_phase:.1f}с, {metrics_count:,} групп")

# ── ЭТАП 3: Salted JOIN со справочником ────────────────────────────────────
print("\n=== Этап 3: Обогащение данными пользователей (Targeted Salting) ===")
t0 = time.time()

# Применяем соль только к горячим ключам
events_with_key = events.withColumn(
    "salt",
    F.when(F.col("user_id").isin(list(hot_keys)), (F.rand() * SALT).cast(IntegerType()))
     .otherwise(F.lit(0))
).withColumn(
    "join_key",
    F.concat_ws("_",
                F.coalesce(F.col("user_id").cast("string"), F.lit("NULL")),
                F.col("salt").cast("string"))
)

# Расширяем users только для горячих ключей
is_hot = F.col("user_id").isin(list(hot_keys))
hot_users = users.filter(is_hot) \
    .withColumn("salt", F.explode(F.array([F.lit(i) for i in range(SALT)])))
normal_users = users.filter(~is_hot).withColumn("salt", F.lit(0))

users_expanded = hot_users.union(normal_users).withColumn(
    "join_key",
    F.concat_ws("_",
                F.coalesce(F.col("user_id").cast("string"), F.lit("NULL")),
                F.col("salt").cast("string"))
)

enriched = events_with_key.join(users_expanded, "join_key", "left") \
    .drop("salt", "join_key") \
    .drop(users_expanded["user_id"])

result_count = enriched.count()
t_join = time.time() - t0
print(f"  Salted JOIN: {t_join:.1f}с, {result_count:,} строк")

# ── Итог ───────────────────────────────────────────────────────────────────
print(f"\n{'='*50}")
print("РЕЗУЛЬТАТ:")
print(f"  Диагностика: {len(hot_keys)} горячих ключей")
print(f"  Two-Phase Aggregation: {t_phase:.1f}с")
print(f"  Targeted Salted JOIN: {t_join:.1f}с")
print(f"  Overhead справочника: {users_expanded.count()}/{users.count()} = "
      f"{users_expanded.count()/users.count():.1f}x (vs {SALT}x при тотальном)")

spark.stop()

Итоги: Salting - оптимизация данных, а не Spark

Salting и двухпроходная агрегация - это инструменты последнего рубежа. Применяйте их только когда исчерпаны более простые решения:

  1. AQE Skew Join (для JOIN-операций с умеренным skew)
  2. Broadcast Join (когда одна сторона достаточно мала)
  3. Repartition с большим числом партиций (временное облегчение)
  4. Bucketing при записи (долгосрочное решение для повторяемых JOIN)

Когда всё это не помогает или неприменимо - salting остаётся единственным способом обработать экстремально скошенные данные без изменения бизнес-логики.

Главное различие между AQE и salting: AQE адаптируется к существующему распределению данных на этапе выполнения. Salting изменяет распределение данных до подачи в Spark. AQE - это умный интерпретатор, salting - это переводчик, который делает данные понятными любому интерпретатору.

Знание обоих подходов и умение выбрать правильный в конкретной ситуации - признак Senior Data Engineer, который понимает не только API, но и физику распределённых вычислений.