sample() и randomSplit(): стратифицированная выборка и train/test split
Bernoulli и Poisson sampling, seed и воспроизводимость, sampleBy для стратификации, randomSplit для ML, ловушки с некэшированными данными и time-based split
Зачем нужна выборка в 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 гарантирует воспроизводимость только если:
- Данные те же самые
- Количество партиций одинаковое
- 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-бюджет команды.