sample() и randomSplit(): стратифицированная выборка и train/test split

Bernoulli и Poisson sampling, seed и воспроизводимость, sampleBy для стратификации, randomSplit для ML, ловушки с некэшированными данными и time-based split

core optimization

Зачем нужна выборка в Data Engineering

С каждым годом объёмы данных в хранилищах растут быстрее, чем вычислительные мощности. Полный датасет в сотни гигабайт или терабайт отлично подходит для ночного batch-джоба, но совершенно неудобен для разработки, отладки и исследовательского анализа. Именно здесь выборка (sampling) становится незаменимым инструментом.

Sampling в контексте Data Engineering и ML решает несколько принципиально разных задач:

Разработка и отладка ETL. Когда вы пишете сложную трансформацию, которая будет работать 4 часа на полном датасете, вы не хотите ждать 4 часа на каждую итерацию. Выборка 1–5% от данных позволяет проверить логику за минуты. Это ускоряет итерации в 20–100 раз.

Исследовательский анализ (EDA). Pandas и визуализация (matplotlib, seaborn) работают только с данными, которые помещаются в оперативную память одной машины. Spark позволяет взять репрезентативную выборку из терабайтного датасета и работать с ней локально.

Подготовка данных для ML. Алгоритмы машинного обучения требуют разделения данных на train/test/validation. Это необходимо делать правильно - воспроизводимо, без утечки данных между наборами, с сохранением распределения классов.

Балансировка датасетов. В задачах обнаружения мошенничества, аномалий, медицинской диагностики положительных примеров может быть 0.1% от всех данных. Обычный random sample уничтожит их. Стратифицированная выборка решает эту проблему.

Приближённые запросы. Иногда точность в 99% лучше чем ответ через 10 минут. Sample-based приближённые агрегаты дают результат в секунды с известной погрешностью.

Как Spark выполняет sampling

Прежде чем переходить к API, важно понять, что происходит внутри Spark при выборке - это объяснит многие неочевидные аспекты поведения функций.

Distributed randomness

В отличие от single-machine библиотек (pandas, scikit-learn), где генератор случайных чисел - один глобальный объект, в Spark каждый Executor работает со своей копией данных и своим локальным генератором случайных чисел. Seed передаётся в задачу и комбинируется с номером partition, чтобы каждая partition генерировала уникальную, но детерминированную последовательность:

Такой подход имеет важное следствие: один и тот же seed с одинаковым числом и распределением partition даёт детерминированный результат. Но если количество partition изменится (например, после repartition или shuffle), seed не поможет - состав выборки изменится, потому что строки попадут в другие partitions.

Bernoulli sampling (без возвращения)

Это стандартный режим sample(withReplacement=False). Spark реализует его через алгоритм Бернулли: для каждой строки независимо генерируется случайное число от 0.0 до 1.0, и если оно меньше заданной доли (fraction), строка включается в выборку.

Математически: каждая строка включается в выборку с вероятностью fraction, независимо от остальных строк. Это означает, что ожидаемое количество строк равно N × fraction, но фактическое количество - случайная величина с биномиальным распределением.

Пример: датасет из 1000 строк, fraction=0.1
  Ожидаемый размер выборки: 100 строк
  Реальный размер: около 100 строк, но может быть 85 или 115

  Стандартное отклонение: sqrt(N × p × (1-p)) = sqrt(1000 × 0.1 × 0.9) ≈ 9.5

  95% случаев: от 81 до 119 строк (100 ± 2σ)

Почему именно такой алгоритм? Потому что он не требует перетасовки всех данных: каждая строка принимает решение о себе сама, локально, без координации с другими Executor'ами. Это делает Bernoulli sampling очень быстрым - по скорости он близок к полному чтению датасета с минимальными дополнительными затратами.

Poisson sampling (с возвращением)

Режим sample(withReplacement=True) использует Poisson sampling: каждая строка может попасть в выборку несколько раз. Это реализуется через распределение Пуассона: для каждой строки генерируется случайное число k из Poisson-распределения с параметром λ = fraction, и строка дублируется k раз (0 раз = не включается).

Bootstrap sampling - это частный случай Poisson sampling с fraction=1.0: каждый элемент может попасть 0, 1, 2 или более раз, и ожидаемый размер выборки равен исходному датасету. Это стандартная техника для оценки дисперсии моделей и confidence intervals в статистике.

from pyspark.sql import SparkSession

spark = SparkSession.builder.appName("sampling-demo").getOrCreate()

# Bootstrap выборка: некоторые строки будут дублированы
df = spark.range(100).withColumnRenamed("id", "row_id")

bootstrap = df.sample(withReplacement=True, fraction=1.0, seed=42)

print(f"Исходный датасет: {df.count()} строк")
print(f"Bootstrap выборка: {bootstrap.count()} строк")  # около 100, но не ровно

# Проверим дублирование: какие строки встречаются более одного раза
from pyspark.sql.functions import count, col

bootstrap.groupBy("row_id") \
    .agg(count("*").alias("occurrences")) \
    .filter(col("occurrences") > 1) \
    .orderBy("row_id") \
    .show(10)
# Некоторые row_id будут встречаться 2-3 раза, другие не появятся вовсе

sample(): базовая вероятностная выборка

Параметры и поведение

# Полная сигнатура
DataFrame.sample(
    withReplacement=False,  # True = Poisson, False = Bernoulli
    fraction=None,          # доля строк (0.0 до 1.0)
    seed=None               # seed для воспроизводимости (None = случайный)
)

Рассмотрим каждый параметр подробнее:

import time
from pyspark.sql import functions as F

# Создадим датасет для экспериментов
transactions = spark.range(1_000_000) \
    .withColumn("user_id",   (col("id") % 10000).cast("int")) \
    .withColumn("amount",    (F.rand() * 1000).cast("double")) \
    .withColumn("category",  F.when(col("id") % 10 == 0, "fraud")
                               .otherwise("normal"))

print(f"Полный датасет: {transactions.count():,} строк")
# Полный датасет: 1,000,000 строк

# Базовая выборка: ~10% строк
sample_10pct = transactions.sample(fraction=0.1, seed=42)
print(f"10% выборка: {sample_10pct.count():,} строк")
# 10% выборка: ~100,000 строк (не ровно 100,000!)

Почему не ровно 10%? Именно из-за Bernoulli-природы sampling: каждая строка независимо бросает монету с вероятностью 0.1. На миллионе строк отклонение составит порядка ±0.3% (несколько тысяч строк), что обычно приемлемо. Но на маленьких датасетах (1000 строк) sample(0.1) даст от 80 до 120 строк - это существенный разброс.

# Демонстрация нестабильности на маленьком датасете
small_df = spark.range(1000)

sizes = []
for i in range(5):
    size = small_df.sample(fraction=0.1, seed=i).count()
    sizes.append(size)
    print(f"  seed={i}: {size} строк")

# seed=0: 94 строки
# seed=1: 107 строк
# seed=2: 101 строка
# seed=3: 96 строк
# seed=4: 98 строк

На маленьких датасетах дисперсия высокая. Если вам нужно ровно N строк - используйте limit(N), а не sample(fraction).

Seed и воспроизводимость

Seed - это инициализирующее значение генератора псевдослучайных чисел. Одинаковый seed на одинаковых данных с одинаковой partition структурой всегда даёт одинаковый результат. Это критически важно для:

  • Воспроизводимых экспериментов: коллеги должны получать тот же train/test split
  • Отладки: если в модели проблема, нужно изолировать её от случайности выборки
  • Compliance: в финансах и медицине эксперименты должны быть полностью воспроизводимы
# Без seed: каждый запуск даёт другую выборку
sample_a = transactions.sample(fraction=0.01)  # seed не задан
sample_b = transactions.sample(fraction=0.01)  # другой результат!

print(transactions.count())
# Пересечение двух выборок без seed - почти случайно

# С seed: детерминированный результат
sample_1 = transactions.sample(fraction=0.01, seed=42)
sample_2 = transactions.sample(fraction=0.01, seed=42)

# Количество совпадает при одинаковых partition
print(f"Выборка 1: {sample_1.count()} строк")
print(f"Выборка 2: {sample_2.count()} строк")  # то же число

# Проверка идентичности: join и сравнение count
joined = sample_1.join(sample_2, on=["id"], how="inner")
print(f"Совпадающих строк: {joined.count()}")
# Совпадающих строк: ~10000 (все)

Важная оговорка о детерминизме. Seed гарантирует воспроизводимость только если:

  1. Данные те же самые
  2. Количество партиций одинаковое
  3. Spark-версия та же (алгоритм sampling может меняться между версиями)

Если между двумя запусками произошёл repartition, shuffle, или изменилось spark.sql.shuffle.partitions - состав выборки изменится даже с тем же seed.

sample() vs limit() vs take(): в чём разница

Это частая точка путаницы для начинающих:

sample(fraction) limit(N) take(N)
Что возвращает Случайные ~fraction% строк Первые N строк логически Первые N строк как Python list
Распределение Репрезентативное (случайное) Смещённое (первые partition) Смещённое (первые partition)
Детерминизм Да (с seed) Да Да
Shuffle Нет Нет Нет
Объём Доля от всего датасета Фиксированное N Фиксированное N
Применение EDA, ML prep Отладка трансформаций Быстрый просмотр
# limit() берёт первые N строк из первых partition - данные будут смещены!
# Если данные отсортированы по времени, limit() вернёт только старые события
first_100 = transactions.limit(100)

# sample() вернёт случайные строки со всего датасета
random_100 = transactions.sample(fraction=0.0001, seed=42)

# Для EDA всегда используйте sample(), не limit()
eda_sample = transactions.sample(fraction=0.01, seed=42).toPandas()
# Теперь можно работать в Pandas: строить гистограммы, корреляции и т.д.

limit() полезен для отладки трансформаций (нужно просто посмотреть на схему и первые записи), но не для статистического анализа.

randomSplit(): разделение на train/test/validation

Как работает randomSplit

randomSplit() принимает список весов (weights) и возвращает список DataFrame'ов. Каждая строка исходного датасета попадает ровно в один из результирующих DataFrame'ов. Под капотом Spark генерирует для каждой строки случайное число и определяет, в какой "сегмент" она попадает:

Важно: randomSplit нормализует веса автоматически. Если передать [80, 20], результат идентичен [0.8, 0.2]. Строка попадает ровно в один сплит, поэтому пересечение множеств всегда пустое.

# Двусторонний split: train/test
train_df, test_df = transactions.randomSplit([0.8, 0.2], seed=42)

print(f"Train: {train_df.count():,} строк")   # ~800,000
print(f"Test:  {test_df.count():,} строк")    # ~200,000

# Трёхсторонний split: train/validation/test
train_df, val_df, test_df = transactions.randomSplit([0.7, 0.15, 0.15], seed=42)

print(f"Train:      {train_df.count():,} строк")  # ~700,000
print(f"Validation: {val_df.count():,} строк")    # ~150,000
print(f"Test:       {test_df.count():,} строк")   # ~150,000

# Проверка отсутствия пересечения (утечки данных)
intersection = train_df.join(test_df, on=["id"], how="inner")
print(f"Пересечение train и test: {intersection.count()} строк")
# Пересечение train и test: 0 строк ✅

Главная ловушка: некэшированный DataFrame

Это самая опасная особенность randomSplit() и самая распространённая ошибка. Если DataFrame не закэширован, Spark пересчитывает его при каждом обращении. Поскольку randomSplit() возвращает несколько независимых DataFrame'ов, каждый из них при материализации запускает отдельный пересчёт исходного DataFrame.

И вот здесь возникает проблема: если источник данных нестабилен (например, это поток Kafka, файловая система с меняющимися файлами, или датасет с недетерминированными трансформациями), при каждом пересчёте строки будут немного другими. В результате одна и та же строка может попасть и в train, и в test - это называется data leakage (утечка данных).

# ❌ НЕПРАВИЛЬНО: DataFrame читается дважды
# При первом count() для train_df исходный файл читается раз
# При первом count() для test_df исходный файл читается снова
# Если файл изменился или строки нестабильны - split будет непоследовательным

train_df, test_df = transactions.randomSplit([0.8, 0.2], seed=42)
print(train_df.count())  # ← первое чтение источника
print(test_df.count())   # ← второе чтение источника (может отличаться!)

# ✅ ПРАВИЛЬНО: кэшировать перед split
transactions.cache()
transactions.count()  # материализуем кэш

train_df, test_df = transactions.cache().randomSplit([0.8, 0.2], seed=42)
print(train_df.count())  # читает из кэша
print(test_df.count())   # читает из кэша - тот же датасет!

# Не забыть освободить кэш после использования
transactions.unpersist()

Почему это работает? Когда DataFrame закэширован, все его обращения читают одни и те же данные из памяти executor'ов. Нет повторного чтения из источника - нет нестабильности.

Практическое правило: перед любым randomSplit() всегда вызывайте .cache() и материализуйте его через .count(). Это немного тратит память, но защищает от неочевидных багов в ML pipeline.

Воспроизводимость и seed

В контексте ML-пайплайнов seed - это не просто удобство, а требование. Без фиксации seed:

  • Коллега не сможет воспроизвести ваши результаты метрик
  • A/B-тест между двумя версиями модели использует разные данные → неверные выводы
  • Сравнение двух алгоритмов нечестное - они могут обучаться на разных подмножествах
# Зафиксируйте seed как константу в начале notebook/скрипта
RANDOM_SEED = 42  # или любое другое число

# Используйте одинаковый seed везде: в split, sample, инициализации модели
train_df, test_df = data.cache().randomSplit([0.8, 0.2], seed=RANDOM_SEED)

# При обучении модели тоже фиксируйте seed
from sklearn.ensemble import RandomForestClassifier
model = RandomForestClassifier(random_state=RANDOM_SEED)

sampleBy(): стратифицированная выборка

Проблема дисбаланса классов

Простой sample() берёт строки равновероятно - каждая строка имеет одинаковый шанс попасть в выборку. Это хорошо работает когда все категории представлены равномерно. Но в реальных задачах это редкость.

Представьте датасет транзакций, где мошеннических операций 0.1% (1000 из 1,000,000). Если взять sample(fraction=0.01), вы ожидаете получить 10,000 строк. Из них мошеннических будет ~10 (0.1% × 10,000). Это критически мало для обучения модели - большинство алгоритмов просто предскажут "нормальная транзакция" для всех строк и получат 99.9% точности, при этом ни разу не поймав мошенника.

Стратифицированная выборка решает это: для каждого класса задаётся своя доля.

sampleBy() API

from pyspark.sql import functions as F

# Создадим датасет с сильным дисбалансом
imbalanced_df = spark.range(1_000_000) \
    .withColumn("label",
        F.when(col("id") % 1000 == 0, 1)   # 0.1% мошенничество
         .otherwise(0)
    ) \
    .withColumn("amount", (F.rand() * 10000).cast("double")) \
    .withColumn("user_id", (col("id") % 50000).cast("int"))

# Проверим распределение
imbalanced_df.groupBy("label").count().show()
# +-----+--------+
# |label|count   |
# +-----+--------+
# |0    |999000  |   ← 99.9% нормальных
# |1    |1000    |   ← 0.1% мошенничество
# +-----+--------+

# sampleBy: задать разные доли для каждого значения ключевой колонки
# fractions - словарь {значение_ключа: доля}
stratified = imbalanced_df.sampleBy(
    col="label",
    fractions={
        0: 0.001,   # взять 0.1% нормальных → ~999 строк
        1: 1.0,     # взять 100% мошеннических → ~1000 строк
    },
    seed=42
)

stratified.groupBy("label").count().show()
# +-----+------+
# |label|count |
# +-----+------+
# |0    |~999  |   ← ~50% нормальных
# |1    |~1000 |   ← ~50% мошеннических
# +-----+------+

print(f"Итого в сбалансированном датасете: {stratified.count():,} строк")
# Итого в сбалансированном датасете: ~1,999 строк

Обратите внимание: значение 1.0 не значит "взять все строки точно", это доля в Bernoulli-sense - каждая строка берётся с вероятностью 1.0, то есть практически все строки попадают в выборку (могут быть минимальные потери из-за статистической природы алгоритма).

Стратификация для многоклассовой классификации

sampleBy() работает не только для бинарной классификации:

# Датасет с несколькими категориями, каждая с разным размером
products = spark.range(500_000) \
    .withColumn("category",
        F.when(col("id") % 100 < 70,  "electronics")   # 70%
         .when(col("id") % 100 < 90,  "clothing")      # 20%
         .when(col("id") % 100 < 99,  "food")          # 9%
         .otherwise(                   "luxury")        # 1%
    )

products.groupBy("category").count().orderBy("count").show()
# +----------+-------+
# |category  |count  |
# +----------+-------+
# |luxury    |5000   |
# |food      |45000  |
# |clothing  |100000 |
# |electronics|350000|
# +----------+-------+

# Стратифицированная выборка: выровнять все категории до ~5000 строк
stratified_products = products.sampleBy(
    col="category",
    fractions={
        "electronics": 5000 / 350000,  # ~0.0143
        "clothing":    5000 / 100000,  # ~0.05
        "food":        5000 / 45000,   # ~0.111
        "luxury":      1.0,            # взять все (и так мало)
    },
    seed=42
)

stratified_products.groupBy("category").count().orderBy("category").show()
# Примерно одинаковое число строк в каждой категории

sampleBy vs sample + filter: почему не filter?

Можно было бы реализовать то же самое через фильтрацию и union:

# Наивная версия без sampleBy (не делайте так):
normal_sample = imbalanced_df.filter(col("label") == 0).sample(fraction=0.001, seed=42)
fraud_sample  = imbalanced_df.filter(col("label") == 1).sample(fraction=1.0,   seed=42)
balanced = normal_sample.union(fraud_sample)

Проблема: Spark читает исходный датасет дважды - по одному разу для каждого filter. sampleBy() читает датасет один раз и принимает решение для каждой строки по её значению ключа. При большом датасете разница в I/O может быть двукратной.

Продвинутые паттерны

Entity-aware split: почему строковый split опасен

Представьте задачу предсказания следующего действия пользователя. У вас есть логи кликов: один пользователь может встречаться в тысячах строк. Если разбить по строкам случайно, события одного пользователя попадут и в train, и в test. Модель обучится на ранних событиях пользователя и будет предсказывать его поздние - это data leakage, завышающий метрики.

# ❌ Строковый split: пользователь 42 в train И в test
train_df, test_df = clicks.randomSplit([0.8, 0.2], seed=42)
# user_id=42 встречается в обоих - утечка!

# ✅ Entity-aware split: разбивать по пользователям, не по строкам
from pyspark.sql.functions import col, hash as spark_hash

# Получаем уникальных пользователей
users = clicks.select("user_id").distinct()

# Делим пользователей (не строки!) на train и test
train_users, test_users = users.randomSplit([0.8, 0.2], seed=42)

# Используем как фильтр для строк
train_df = clicks.join(train_users, on="user_id", how="inner")
test_df  = clicks.join(test_users,  on="user_id", how="inner")

# Проверяем: ни один пользователь не в обоих наборах
train_user_ids = train_users.select("user_id")
test_user_ids  = test_users.select("user_id")
overlap = train_user_ids.join(test_user_ids, on="user_id", how="inner")
print(f"Пересечение пользователей: {overlap.count()} (должно быть 0)")

Это более медленный подход (лишний join), но он единственный корректный для задач с entity-based данными.

Time-based split: для временных рядов

В задачах прогнозирования на временных рядах (цены, трафик, поведение пользователей) случайный split принципиально неверен. Модель не должна "знать будущее" при обучении. Train должен содержать прошлое, test - будущее:

from pyspark.sql.functions import col, to_date, lit

# Данные с колонкой event_date
events = spark.createDataFrame([
    ("EVT-001", "2024-01-15", 42, 1),
    ("EVT-002", "2024-06-20", 17, 0),
    ("EVT-003", "2024-11-05", 99, 1),
    ("EVT-004", "2024-12-01", 42, 0),
], ["event_id", "event_date", "user_id", "label"]) \
    .withColumn("event_date", to_date(col("event_date")))

# Определяем дату разреза (80% данных по времени)
split_date = "2024-10-01"  # первые 10 месяцев - train

train_df = events.filter(col("event_date") <  lit(split_date))
test_df  = events.filter(col("event_date") >= lit(split_date))

print(f"Train: {train_df.count()} событий до {split_date}")
print(f"Test:  {test_df.count()} событий после {split_date}")

# Time-based split не требует seed - граница детерминирована датой
# Но требует хорошего выбора split_date: обычно берут последние X% времени

Time-based split - единственный корректный подход для задач, где важна темпоральность: предсказание транзакций, спроса, отказов оборудования, поведения пользователей.

Sampling для анализа перекоса данных (skew analysis)

Spark внутренне использует sampling при выполнении Sort-Merge Join для оценки распределения ключей (это основа алгоритма Range Partitioner). Вы можете использовать тот же принцип для диагностики data skew перед запуском тяжёлых операций:

# Вместо того чтобы делать groupBy на 1 ТБ датасете,
# сначала посмотрим на распределение ключей на выборке

# Оцениваем skew через sample
key_distribution = (
    huge_df
    .sample(fraction=0.01, seed=42)   # 1% - достаточно для оценки
    .groupBy("join_key")
    .count()
    .orderBy(col("count").desc())
)

key_distribution.show(20)
# Если топ-1 ключ встречается в 50% строк выборки - это сильный skew
# Можно принять решение: использовать salting или AQE skew join

# Оценка процента данных для топ-10 ключей
total_sample = key_distribution.agg(F.sum("count").alias("total")).collect()[0]["total"]
top10 = key_distribution.limit(10).agg(F.sum("count").alias("top10_count")).collect()[0]["top10_count"]
print(f"Топ-10 ключей содержат {100 * top10 / total_sample:.1f}% строк в выборке")

Это практика "сначала исследуй, потом оптимизируй": по выборке 1% можно за секунды понять распределение ключей, которое на полном датасете занимало бы минуты.

Sampling для ускорения разработки ETL

Наиболее часто применяемый паттерн в ежедневной работе DE:

# Паттерн: разрабатывать на sample, запускать на полных данных

DEV_MODE = True   # переключатель

# Читаем данные
raw_data = spark.read.parquet("/data/warehouse/events/")

# Если разработка - берём 1%
if DEV_MODE:
    raw_data = raw_data.sample(fraction=0.01, seed=42).cache()
    raw_data.count()  # материализуем
    print(f"DEV MODE: работаем с {raw_data.count():,} строками")
else:
    print(f"PROD MODE: работаем с {raw_data.count():,} строками")

# Дальше вся логика трансформации одинакова для обоих режимов
result = (
    raw_data
    .filter(col("event_type") == "purchase")
    .groupBy("user_id", "product_category")
    .agg(
        F.sum("amount").alias("total_spent"),
        F.count("*").alias("purchase_count"),
    )
)

if DEV_MODE:
    result.show(20)
else:
    result.write.mode("overwrite").parquet("/data/warehouse/mart/user_purchases/")

Такой паттерн экономит часы разработки: сначала вся логика отлаживается на 1% за минуты, затем переключатель меняется на False и джоб запускается на полных данных.

Explain plan для sampling операций

Полезно понимать, как Spark включает sampling в физический план выполнения:

transactions.sample(fraction=0.1, seed=42).explain("formatted")
# == Physical Plan ==
# Sample 0.1, 42, false   ← fraction, seed, withReplacement
# +- *(1) Project [id#0L, user_id#1, amount#2, category#3]
#    +- *(1) Range (0, 1000000, step=1, splits=8)

# Для randomSplit
train_df, test_df = transactions.randomSplit([0.8, 0.2], seed=42)
train_df.explain("formatted")
# == Physical Plan ==
# Sample 0.0, 0.8, 42, false   ← нижняя граница, верхняя граница, seed
# +- *(1) Project [...]

Что важно: Sample оператор появляется прямо над источником данных. Это означает, что Spark применяет фильтрацию на этапе чтения - не нужно читать все данные в память и потом фильтровать. Sample operator встраивается в plan до materialisation. По этой же причине sample на columnar форматах (Parquet) всё равно читает все столбцы - column pruning работает отдельно.

Практика: подготовка данных для ML-пайплайна

Разберём полный сценарий: таблица кликов рекламных объявлений, задача - предсказать будет ли клик (CTR prediction).

from pyspark.sql import SparkSession
from pyspark.sql import functions as F
from pyspark.sql.types import *

spark = SparkSession.builder.appName("ml-sampling-demo").getOrCreate()

# Симулируем реальные данные рекламного аукциона
# 10 миллионов показов, CTR ~1% (100k кликов)
N = 10_000_000
ads_data = spark.range(N) \
    .withColumn("user_id",     (col("id") % 500_000).cast("int")) \
    .withColumn("ad_id",       (col("id") % 50_000).cast("int")) \
    .withColumn("platform",    F.when(col("id") % 3 == 0, "ios")
                                .when(col("id") % 3 == 1, "android")
                                .otherwise("web")) \
    .withColumn("hour",        (col("id") % 24).cast("int")) \
    .withColumn("clicked",
        F.when(
            (col("id") % 100 == 0) |           # базовые клики ~1%
            ((col("hour") >= 18) & (col("id") % 50 == 0)),  # вечерний буст
            1
        ).otherwise(0)
    ) \
    .withColumn("event_date",
        F.date_add(F.lit("2024-01-01"), (col("id") % 90).cast("int"))
    )

print(f"Всего записей: {ads_data.count():,}")

# Проверим CTR
ads_data.groupBy("clicked").count() \
    .withColumn("pct", F.round(col("count") / N * 100, 2)) \
    .show()
# +-------+--------+----+
# |clicked|count   |pct |
# +-------+--------+----+
# |0      |9780000 |97.8|   ← 97.8% показов без клика
# |1      |220000  |2.2 |   ← 2.2% кликов
# +-------+--------+----+

# ── Шаг 1: Быстрый EDA на выборке ─────────────────────────────────────────────
# Для визуализации в Pandas нужно не более 50-100k строк
eda_sample = ads_data.sample(fraction=0.005, seed=42).cache()
eda_sample.count()  # материализуем

print(f"EDA выборка: {eda_sample.count():,} строк")

# Работаем с Pandas для EDA
eda_pandas = eda_sample.toPandas()
print(eda_pandas.describe())
print(f"\nCTR в выборке: {eda_pandas['clicked'].mean():.3f}")
# CTR в выборке: ~0.022 - близко к реальным 2.2%

eda_sample.unpersist()

# ── Шаг 2: Стратифицированная выборка для обучения модели ─────────────────────
# Хотим: ~100k строк с кликом, ~100k без клика (50/50)
# Кликов: 220k → берём 100k/220k ≈ 45.5%
# Без кликов: 9.78M → берём 100k/9.78M ≈ 1.02%

total_clicked    = 220_000   # примерное число кликов
total_no_clicked = 9_780_000  # примерное число без кликов
target_per_class = 100_000

balanced = ads_data.sampleBy(
    col="clicked",
    fractions={
        0: target_per_class / total_no_clicked,  # ~0.0102
        1: target_per_class / total_clicked,     # ~0.4545
    },
    seed=42
)

# Кэшируем - будем обращаться несколько раз
balanced.cache()
balanced.count()  # материализуем

# Проверяем баланс
balanced.groupBy("clicked").count().show()
# +-------+-------+
# |clicked|count  |
# +-------+-------+
# |0      |~100000|   ← ~50%
# |1      |~100000|   ← ~50%
# +-------+-------+

# ── Шаг 3: Train/Test split ───────────────────────────────────────────────────
# DataFrame уже кэширован - безопасно делать randomSplit
train_df, test_df = balanced.randomSplit([0.8, 0.2], seed=42)

print(f"Train: {train_df.count():,} строк")
print(f"Test:  {test_df.count():,} строк")

# ── Шаг 4: Проверка распределения в обоих наборах ─────────────────────────────
# Убеждаемся что stратификация сохранилась в train и test

def check_class_balance(df, name):
    total = df.count()
    dist = df.groupBy("clicked").count() \
              .withColumn("pct", F.round(col("count") / total * 100, 1))
    print(f"\n{name} ({total:,} строк):")
    dist.show()

check_class_balance(train_df, "Train")
check_class_balance(test_df,  "Test")

# Train (160,000 строк):
# +-------+-------+----+
# |clicked|count  |pct |
# +-------+-------+----+
# |0      |80000  |50.0|
# |1      |80000  |50.0|
# +-------+-------+----+

# Test (40,000 строк):
# +-------+------+----+
# |clicked|count |pct |
# +-------+------+----+
# |0      |20000 |50.0|
# |1      |20000 |50.0|
# +-------+------+----+

# Проверяем отсутствие data leakage
overlap = train_df.join(test_df, on=["id"], how="inner")
print(f"\nData leakage (пересечение строк): {overlap.count()} строк")
# Data leakage: 0 строк ✅

# ── Шаг 5: Освобождаем кэш ────────────────────────────────────────────────────
balanced.unpersist()

Anti-patterns и частые ошибки

1. randomSplit без кэша. Вызов randomSplit() на некэшированном DataFrame, который затем используется в нескольких Action - гарантированный источник неожиданного поведения.

2. sample() после тяжёлого shuffle. Если сделать groupBy, затем sample, Spark сначала выполнит весь shuffle (дорого), потом отфильтрует строки. Правильно: sample сначала, потом groupBy - Spark уменьшит объём данных до тяжёлой операции.

# ❌ Дорого: сначала shuffle на всех данных, потом sample
result = huge_df.groupBy("category").count().sample(fraction=0.1, seed=42)

# ✅ Дешевле: сначала уменьшить объём, потом группировать
# (только если вы делаете exploratory aggregation, не точную)
result = huge_df.sample(fraction=0.1, seed=42).groupBy("category").count()

3. Использование limit() вместо sample() для EDA. limit(1000) вернёт первые 1000 строк из первых partition - это смещённая, неслучайная выборка. В Parquet это могут быть все данные за один день. Для исследования используйте sample().

4. Игнорирование entity leakage. Разбивка по строкам для данных с identity (пользователи, устройства, сессии) создаёт утечку данных и завышает метрики модели. Всегда анализируйте структуру данных перед split.

5. Разные seed в разных частях пайплайна. Если sample использует seed=42, а model training использует random_state=100, то при смене seed в одном месте другая часть остаётся с "прежним" random - результаты несовместимы. Используйте одну константу RANDOM_SEED на весь пайплайн.

Best Practices

Всегда фиксируй seed. Нет seed - нет воспроизводимости. Даже в development-скриптах: сегодняшний dev-скрипт - это завтрашний production-джоб.

Кэшируй перед randomSplit(). Вызов .cache() и материализация через .count() перед split - обязательный паттерн. Без этого split работает нестабильно на неидемпотентных источниках.

Проверяй баланс классов после split. После stratified sampling и randomSplit проверьте что распределение целевой переменной одинаково в train и test. Это занимает 2 строки кода и может предотвратить часы отладки.

Для entity-данных разбивай по entity, не по строкам. Пользователи, устройства, сессии, документы - всегда делай split по уникальным entity.

Для временных рядов используй time-based split. Случайный split на temporal данных - это предсказание прошлого по будущему. Граница должна быть датой, а не случайным числом.

Используй sample для диагностики skew. Перед тяжёлым join проверь распределение ключей на 1% выборке - это секунды вместо минут.

Развертывай ETL на sample сначала. Паттерн DEV_MODE = True/False с sample(0.01) ускоряет разработку в 10–100 раз и экономит compute-бюджет команды.