Actions в деталях: foreachPartition, checkpoint для длинных DAG, toPandas safety
Анатомия Spark Actions: DAG scheduling, foreach vs foreachPartition, connection pool, checkpoint vs cache, toPandas OOM и Arrow-оптимизация
Почему Actions - самое дорогое место в Spark¶
Spark работает по модели ленивых вычислений (lazy evaluation). Любой вызов .filter(), .join(), .groupBy() - не выполняется немедленно. Spark только строит граф вычислений (DAG) в памяти Driver. Реальные вычисления начинаются в одной-единственной точке: при вызове Action.
Action - это команда «теперь действительно посчитай». После её вызова запускается полная цепочка: Catalyst оптимизирует план, DAG Scheduler разбивает его на Stage, Task Scheduler раздаёт Tasks по Executor-ам, данные перемещаются по сети, результат собирается на Driver.
Именно поэтому Action - самая дорогая операция в Spark. Разработчики часто не замечают скрытых расходов: лишний count() в цикле, collect() который стягивает гигабайты на Driver, foreach() который открывает тысячи соединений к базе данных. Всё это - антипаттерны, которые превращают высокопроизводительный распределённый движок в медленный, а порой и аварийно завершающийся процесс.
В этом уроке разберём механику выполнения Action изнутри, опасные паттерны и способы их исправления.
Что происходит при вызове Action¶
Вызов Action запускает многоуровневую цепочку событий, большая часть которых происходит внутри Driver-процесса.
Фазы выполнения¶
Фаза 1: Catalyst Optimization. Driver берёт накопленный Logical Plan (граф трансформаций) и прогоняет его через Catalyst Optimizer. Catalyst применяет правила: предикатный pushdown (перемещение фильтров ближе к источнику данных), column pruning (удаление лишних колонок), constant folding (вычисление константных выражений). На выходе - оптимизированный Physical Plan.
Фаза 2: DAG Scheduling. DAG Scheduler анализирует Physical Plan и разбивает его на Stages (стадии). Граница между стадиями - это Shuffle (перераспределение данных по сети). Внутри одной стадии данные не перемещаются между executor-ами. Каждая стадия состоит из Tasks - одна Task обрабатывает одну партицию данных.
Фаза 3: Task Scheduling. Task Scheduler назначает Tasks на свободные Executor-ы с учётом локальности данных (data locality): предпочтительно запускать Task на том executor-е, где физически находятся нужные данные - это экономит сетевой трафик.
Фаза 4: Physical Execution. Executor-ы выполняют Tasks параллельно. Каждая Task обрабатывает свою партицию: читает данные, применяет трансформации, записывает промежуточные результаты при shuffle или финальные - в sink (файл, таблица, внешняя система).
Фаза 5: Result Collection. Результат Tasks (для collect, count, take) или подтверждение записи (для write, foreach) возвращается на Driver. Driver собирает частичные результаты и возвращает итог в пользовательский код.
Диаграмма показывает полный путь от вызова Action до получения результата. Ключевой момент: если план содержит shuffle (например, groupBy, join, repartition), он разбивается на несколько Stage. Каждый переход между Stage - это барьер: Spark ждёт завершения всех Tasks предыдущей Stage перед началом следующей.
Почему даже count() может быть дорогим¶
df.count() выглядит безобидно, но за ним стоит полное выполнение всего DAG вплоть до финальной агрегации. Если DataFrame получен после нескольких join и filter, каждый count() заново запускает весь этот граф. Без кэширования три count() в коде = три полных прогона pipeline.
Каталог Actions: классификация по назначению¶
Все Actions делятся на четыре категории по типу результата и назначению.
Агрегирующие (возвращают скаляр или коллекцию скаляров):
count()- количество строк; запускает полный DAG с финальной агрегацией count на Driversum(col),max(col),min(col)- через.agg()или SQL; также полный DAGfirst(),head(n)- возвращает первую строку или n строк; может использовать оптимизацию ранней остановки
Сбор данных на Driver:
collect()- все строки →List[Row]в памяти Driver; опасен при больших данныхtake(n)- первые n строк; аналогичноcollect(), но ограниченtoPandas()- все строки →pandas.DataFrameв памяти Driver; обсудим подробноtoLocalIterator()- lazy итератор по строкам; не стягивает всё сразу, но последовательный
Запись во внешние хранилища:
write.csv(),write.parquet(),write.saveAsTable()- запись в файловую систему или таблицуwriteStream.start()- запуск Structured Streaming
Пользовательская обработка:
foreach(f)- вызываетfдля каждой строки; запускает f на executor-ахforeachPartition(f)- вызываетfдля каждой партиции (Iterator[Row]); более эффективно
Правило: Actions первых двух категорий (агрегирование и сбор) стягивают данные на Driver - это bottleneck. Actions четвёртой категории выполняются на executor-ах и не загружают Driver - это правильная архитектура для вывода данных во внешние системы.
collect() и driver bottleneck¶
collect() - самый прямолинейный Action: все строки DataFrame стягиваются на Driver и возвращаются как Python-список объектов Row.
# Антипаттерн: collect() на большом датасете
all_rows = df.collect() # ← может потребовать gigabytes памяти на Driver
for row in all_rows:
process(row) # последовательная обработка на Driver
# Правильно: обработка остаётся на Executors
df.foreachPartition(lambda rows: [process(row) for row in rows])
Проблемы collect() в production:
- Driver OOM: если датасет больше доступной памяти Driver (
spark.driver.memory), процесс упадёт сjava.lang.OutOfMemoryError. Driver по умолчанию имеет 1–4 ГБ, тогда как данные могут быть в сотни ГБ - Сетевой трафик: все данные со всех executor-ов передаются по сети на Driver. При 100 ГБ данных это 100 ГБ трафика - медленно и дорого
- Сериализация: каждый
Rowсериализуется на executor-е и десериализуется на Driver. Без Arrow - это pickle, с миллионами строк - CPU bottleneck
Безопасное использование collect():
collect() допустим только когда данные заведомо малы (результат агрегации, lookup-таблицы, выборка):
# OK: после агрегации результат мал
metrics = df.groupBy("region").agg({"revenue": "sum"}).collect()
# OK: небольшая lookup-таблица для broadcast join
lookup_rows = spark.table("dim_products").filter("is_active = true").collect()
toPandas(): как распределённые данные попадают в один процесс¶
toPandas() - один из самых часто используемых и одновременно наиболее опасных Actions в PySpark. Он создаёт фундаментальный архитектурный барьер: все распределённые данные с множества executor-ов стягиваются в память единственного Driver-процесса и превращаются в pandas.DataFrame.
Механика toPandas()¶
Диаграмма показывает фундаментальную проблему: независимо от размера кластера (10, 100, 1000 executor-ов), все данные в итоге оказываются в одном процессе - Driver. Это полностью аннулирует распределённость вычислений. Если три executor-а держали по 2 ГБ данных, Driver должен принять и разместить в своей куче 6 ГБ.
Анатомия Driver OOM¶
Типичный сценарий: дата-инженер пишет notebook, запускает df.toPandas() для построения графика. Датасет «небольшой» - всего 5 миллионов строк по 50 колонок. Но каждая строка в pandas занимает ~8 байт на числовое поле × 50 колонок = 400 байт × 5M строк = 2 ГБ данных плюс overhead на pandas-структуры. Driver с настройкой spark.driver.memory=2g немедленно падает с OOM.
Спасти ситуацию может Apache Arrow.
Arrow-оптимизация для toPandas()¶
# Включить Arrow для toPandas() и createDataFrame()
spark.conf.set("spark.sql.execution.arrow.pyspark.enabled", "true")
# Теперь toPandas() использует Arrow RecordBatch вместо pickle
df_pandas = df.toPandas()
С arrow.pyspark.enabled = true процесс меняется кардинально:
- Без Arrow: каждая строка сериализуется через pickle на executor-е → передаётся на Driver → десериализуется из pickle → добавляется в pandas DataFrame. Для 10 млн строк - 10 млн операций сериализации/десериализации
- С Arrow: колонка данных конвертируется в Arrow RecordBatch (непрерывный бинарный буфер). Передача - одна операция IO на батч. Driver получает Arrow RecordBatch и конвертирует в pandas через
pyarrow- для числовых данных это почти zero-copy
Практические результаты:
| Метод | Время (5M строк, 10 числовых колонок) | Memory overhead |
|---|---|---|
| Без Arrow (pickle) | ~85 сек | Высокий |
| С Arrow | ~12 сек | Умеренный |
| Arrow + batching | ~10 сек | Минимальный |
Размер батча Arrow настраивается:
# Размер одного Arrow RecordBatch (строк)
spark.conf.set("spark.sql.execution.arrow.maxRecordsPerBatch", "50000")
Большой батч уменьшает число IPC-передач, но требует больше памяти на executor-е для формирования буфера. Для датасетов с широкими схемами (50+ колонок) имеет смысл уменьшить батч.
Безопасные паттерны для toPandas()¶
# Паттерн 1: limit() перед выгрузкой - защита от OOM
# Для визуализации и дебага обычно нужно не более 10 000 строк
df_sample = df.limit(10_000).toPandas()
# Паттерн 2: sample() для статистического анализа
# 0.1% от датасета часто достаточно для EDA
df_sample = df.sample(fraction=0.001, seed=42).toPandas()
# Паттерн 3: агрегация ПЕРЕД выгрузкой
# Вместо выгрузки сырых данных - выгружаем агрегат
df_agg = (
df.groupBy("product_category")
.agg({"revenue": "sum", "orders": "count"})
.orderBy("sum(revenue)", ascending=False)
.limit(50)
.toPandas()
)
# Паттерн 4: выгрузка только нужных колонок
# Не toPandas() на широкий датасет - только нужные поля
df_narrow = df.select("date", "metric_a", "metric_b").limit(5000).toPandas()
# Паттерн 5: проверка размера ПЕРЕД выгрузкой
row_count = df.count()
col_count = len(df.columns)
est_mb = row_count * col_count * 8 / 1024 / 1024 # грубая оценка в МБ
if est_mb > 500:
raise ValueError(
f"DataFrame слишком большой для toPandas(): ~{est_mb:.0f} МБ. "
f"Используйте .limit() или .sample() перед выгрузкой."
)
df_pandas = df.toPandas()
PySpark Pandas API: распределённый альтернативный путь¶
Если вам нужна pandas-подобная работа с данными, но они слишком велики для toPandas(), рассмотрите PySpark Pandas API (бывший Koalas):
import pyspark.pandas as ps
# Создать pandas-on-Spark DataFrame - вычисления остаются распределёнными
psdf = ps.from_pandas(df_pandas) # или ps.read_csv(path)
# API идентичен pandas, но выполняется на Spark кластере
psdf_result = psdf.groupby("region")["revenue"].sum().sort_values(ascending=False)
# Конвертировать обратно в pandas только финальный (маленький) результат
result = psdf_result.to_pandas()
PySpark Pandas API переводит pandas-вызовы в Spark операции под капотом. Вычисления остаются распределёнными - нет переноса всех данных на Driver. Ограничения: не все pandas-функции поддержаны, и некоторые операции (например, iloc) требуют глобальной сортировки, что дорого.
foreach vs foreachPartition: критическое различие¶
foreach(f) и foreachPartition(f) - оба Action, оба выполняют пользовательскую функцию на executor-ах. Но между ними принципиальное архитектурное различие.
foreach: антипаттерн для внешних систем¶
foreach(f) вызывает функцию f для каждой строки датасета по отдельности. Если у вас 100 000 строк, функция вызывается 100 000 раз.
# АНТИПАТТЕРН: foreach для записи в PostgreSQL
def write_row_to_db(row):
conn = psycopg2.connect(CONN_STRING) # открывает соединение на каждую строку!
cursor = conn.cursor()
cursor.execute(
"INSERT INTO events (id, ts, value) VALUES (%s, %s, %s)",
(row["id"], row["ts"], row["value"])
)
conn.commit()
cursor.close()
conn.close() # закрывает соединение - итого 100 000 open/close пар
df.foreach(write_row_to_db)
При 100 000 строках этот код:
- открывает и закрывает 100 000 TCP-соединений к PostgreSQL
- выполняет 100 000 отдельных INSERT-запросов (вместо одного batch INSERT)
- перегружает connection pool PostgreSQL и его WAL
- работает в 100–1000× медленнее, чем batch-вставка
foreachPartition: правильный паттерн¶
foreachPartition(f) вызывает функцию f один раз для каждой партиции, передавая Iterator по строкам партиции. Открыть соединение можно один раз на всю партицию.
На диаграмме видна принципиальная разница: foreach открывает соединение для каждой строки отдельно - это катастрофически неэффективно. foreachPartition открывает одно соединение на всю партицию и использует batch-запросы.
Проблема сериализации: Task not serializable¶
Почему нельзя создать соединение на Driver и передать его в замыкание? Потому что Spark должен сериализовать функцию (вместе с её замыканием) и отправить на каждый executor. Большинство объектов соединений (JDBC Connection, Kafka Producer, HTTP Client) не реализуют Java Serializable - и Spark выбросит исключение.
# ОШИБКА: создаём соединение на Driver - оно не сериализуемо
conn = psycopg2.connect(CONN_STRING) # объект на Driver
def write_row(row):
conn.cursor().execute(...) # ← Spark пытается сериализовать conn → ОШИБКА
df.foreach(write_row)
# SparkException: Task not serializable
Решение: создавать соединение внутри функции, которая выполняется на executor-е:
# ПРАВИЛЬНО: соединение создаётся на executor-е, не передаётся через сеть
def write_partition(rows):
conn = psycopg2.connect(CONN_STRING) # ← внутри функции, создаётся на executor-е
cursor = conn.cursor()
batch = []
try:
for row in rows:
batch.append((row["id"], row["ts"], row["value"]))
if len(batch) >= 1000:
cursor.executemany(
"INSERT INTO events (id, ts, value) VALUES (%s, %s, %s)", batch
)
conn.commit()
batch = []
if batch: # остаток последнего неполного батча
cursor.executemany(
"INSERT INTO events (id, ts, value) VALUES (%s, %s, %s)", batch
)
conn.commit()
finally:
cursor.close()
conn.close()
df.foreachPartition(write_partition)
Connection Pool: продвинутый паттерн¶
Для высоконагруженных pipeline создание нового соединения при каждом вызове write_partition (то есть на каждую партицию) тоже может быть дорогим. Продвинутый паттерн - использование пула соединений на уровне JVM executor-а через статические/глобальные переменные Python.
# Пул соединений на уровне Python-процесса executor-а
# Создаётся один раз при первом вызове на данном executor-е
_connection_pool = None
def get_connection_pool():
"""Ленивая инициализация пула соединений на executor-е."""
global _connection_pool
if _connection_pool is None:
from psycopg2 import pool as pg_pool
_connection_pool = pg_pool.ThreadedConnectionPool(
minconn=1,
maxconn=4, # максимум 4 соединения на executor
dsn=CONN_STRING
)
return _connection_pool
def write_partition_with_pool(rows):
"""Пишет партицию в БД, переиспользуя соединение из пула."""
pool = get_connection_pool()
conn = pool.getconn()
try:
cursor = conn.cursor()
batch = []
for row in rows:
batch.append((row["id"], row["ts"], row["value"]))
if len(batch) >= 1000:
cursor.executemany(INSERT_SQL, batch)
conn.commit()
batch = []
if batch:
cursor.executemany(INSERT_SQL, batch)
conn.commit()
cursor.close()
except Exception:
conn.rollback()
raise
finally:
pool.putconn(conn) # возвращаем в пул, не закрываем
df.foreachPartition(write_partition_with_pool)
Пул создаётся один раз за всё время жизни executor-а (Python-процесса) и переиспользуется для всех партиций, которые обрабатывает этот executor. Это значительно снижает overhead на установку TCP-соединения, SSL-хендшейк и аутентификацию.
Важный нюанс: глобальная переменная _connection_pool живёт в Python-процессе executor-а. При рестарте executor-а (из-за OOM или других ошибок) пул создаётся заново. Это корректное поведение - пул не «протекает» между executor-ами.
Batch writes и идемпотентность¶
При записи через foreachPartition необходимо учитывать семантику at-least-once: если Task упала и была перезапущена, данные могут быть записаны дважды. Spark не гарантирует exactly-once для пользовательских Actions.
def write_partition_idempotent(rows):
"""Идемпотентная запись: при повторе не создаёт дубликатов."""
conn = psycopg2.connect(CONN_STRING)
cursor = conn.cursor()
try:
for row in rows:
# INSERT ... ON CONFLICT DO NOTHING - идемпотентно
cursor.execute(
"""
INSERT INTO events (id, ts, value)
VALUES (%s, %s, %s)
ON CONFLICT (id) DO NOTHING
""",
(row["id"], row["ts"], row["value"])
)
conn.commit()
finally:
cursor.close()
conn.close()
Для аналитических баз данных (ClickHouse, BigQuery) используйте INSERT IGNORE или staging table + merge strategy.
Проблема длинных DAG¶
По мере роста pipeline число трансформаций увеличивается. В production встречаются pipeline из 50-100 последовательных join, union, filter. Каждая трансформация добавляет узел в Logical Plan Catalyst. При достаточно длинном графе возникают серьёзные проблемы.
Симптомы слишком длинного DAG¶
StackOverflowError при планировании. Catalyst Optimizer анализирует граф рекурсивно. Если граф содержит 200+ узлов, рекурсия переполняет стек JVM:
java.lang.StackOverflowError
at org.apache.spark.sql.catalyst.trees.TreeNode.foreach(TreeNode.scala:...)
Медленная оптимизация. Catalyst применяет правила оптимизации итеративно. При длинном плане это занимает секунды, что добавляется к общему времени каждого Action.
Медленный пересчёт при сбое. Если executor упал и Task нужно перезапустить, Spark прослеживает всю lineage (историю вычислений) от источника до текущей точки. Для 100-шаговой lineage - это 100 шагов пересчёта.
Когда это особенно критично:
- итеративные алгоритмы в цикле (PageRank, k-means, алгоритмы на графах)
- динамическое построение запросов через union в цикле
- рекурсивная обработка иерархических данных
- pipeline с 50+ последовательными join
cache() и persist() - кэширование с сохранением lineage¶
cache() / persist() - первый инструмент для повторного использования результатов вычислений. При вызове df.cache() Spark материализует DataFrame в памяти executor-ов (или на диске, в зависимости от StorageLevel). При повторном использовании DataFrame Spark читает данные из кэша, а не перевычисляет.
from pyspark import StorageLevel
# cache() - хранить в памяти (MEMORY_AND_DISK_DESER)
df_cached = df.expensive_join().filter(condition).cache()
df_cached.count() # первый Action материализует кэш
# persist() - явный контроль уровня хранения
df_persisted = df.join(lookup, "id").persist(StorageLevel.MEMORY_AND_DISK)
# ВАЖНО: освобождать кэш когда больше не нужен
df_cached.unpersist()
Но у cache() есть фундаментальное ограничение: lineage сохраняется. Если данные из кэша потеряны (executor упал, кэш вытеснен из памяти), Spark пересчитывает их заново по lineage. Это хорошо для fault tolerance, но для длинных DAG пересчёт может занять очень много времени.
checkpoint(): разрыв lineage¶
checkpoint() - радикальное решение проблемы длинного DAG. В отличие от cache(), checkpoint физически записывает DataFrame на отказоустойчивое хранилище (HDFS, S3, ADLS) и полностью забывает его lineage. После checkpoint Spark не знает, как этот DataFrame был вычислен - он просто читает его с диска.
Диаграмма демонстрирует эффект checkpoint: вместо одной непрерывной цепи из 50 шагов получаем несколько коротких сегментов. При сбое Spark нужно пересчитать не всё с нуля, а только последний сегмент (10 шагов вместо 50). Логический план Catalyst при каждом checkpoint сбрасывается до одного узла «читать с диска».
Настройка и использование checkpoint¶
# Шаг 1: ОБЯЗАТЕЛЬНО - установить директорию для checkpoint ПЕРЕД использованием
spark.sparkContext.setCheckpointDir("s3a://data-lake-checkpoints/spark-checkpoints/")
# Шаг 2: Выполнить checkpoint
# По умолчанию checkpoint() EAGER: немедленно запускает Action для материализации
df_checkpointed = df.expensive_computation().checkpoint()
# Lazy checkpoint: не запускает Action немедленно
# Материализуется при следующем Action
df_lazy = df.expensive_computation().checkpoint(eager=False)
df_lazy.count() # здесь происходит и checkpoint, и count
# Шаг 3: Продолжаем работу - lineage уже разорвана
result = df_checkpointed.further_transformation().action()
Разница между eager и lazy checkpoint важна для управления памятью и временем выполнения:
- Eager (по умолчанию): сразу запускает Job для записи на диск. Логический план Catalyst сразу сбрасывается. Хорошо для предсказуемости, плохо если Action нужен всё равно
- Lazy: записывает на диск при следующем Action - это экономит лишний запуск Job если Action всё равно нужен сразу после
localCheckpoint: скорость vs надёжность¶
localCheckpoint() - вариант, который записывает данные на локальный диск executor-а, а не на распределённое хранилище:
df_local = df.heavy_computation().localCheckpoint()
Преимущества localCheckpoint:
- значительно быстрее, чем checkpoint на S3/HDFS (нет сетевого трансфера)
- разрывает lineage так же, как и checkpoint()
- подходит для промежуточных результатов, которые не нужно восстанавливать после полного сбоя кластера
Недостатки localCheckpoint:
- данные теряются при падении executor-а - Spark не может их восстановить
- если executor-а нет, соответствующие партиции не доступны → Job завершается с ошибкой
- не подходит для production-критичных данных
Когда использовать localCheckpoint: для ускорения итеративных алгоритмов в рамках одного Job, где fault tolerance не критична (тестовые среды, ML-экспериментирование).
checkpoint vs cache/persist: сравнение¶
Диаграмма - дерево решений для выбора между cache, checkpoint и localCheckpoint. Если DAG короткий, достаточно cache(). Если длинный и нужна fault tolerance - checkpoint() на S3/HDFS. Если fault tolerance не критична, но нужна скорость - localCheckpoint().
| Характеристика | cache() |
checkpoint() |
localCheckpoint() |
|---|---|---|---|
| Lineage | Сохраняется | Разрывается | Разрывается |
| Хранилище | Память/диск executor | S3 / HDFS (надёжно) | Локальный диск executor |
| Fault tolerance | Пересчёт по lineage | Чтение с диска | Нет (если executor упал) |
| Скорость записи | Быстро | Медленно (сеть) | Быстро |
| Освобождение памяти | unpersist() |
Ручное удаление с диска | Автоматически |
checkpoint для итеративных алгоритмов¶
Самый важный сценарий для checkpoint - итеративные алгоритмы, где граф растёт с каждой итерацией:
spark.sparkContext.setCheckpointDir("s3a://bucket/checkpoints/")
CHECKPOINT_INTERVAL = 5 # делать checkpoint каждые 5 итераций
# Итеративный алгоритм (например, упрощённый PageRank)
current_df = initial_df.cache()
for iteration in range(50):
# Следующая итерация - новая трансформация добавляется к графу
next_df = compute_next_iteration(current_df, adjacency_df)
if (iteration + 1) % CHECKPOINT_INTERVAL == 0:
# Без checkpoint DAG вырастет до 50 вложенных compute_next_iteration()
# С checkpoint - граф сбрасывается, следующая итерация читает с диска
next_df = next_df.checkpoint()
current_df.unpersist() # освобождаем предыдущий кэш
current_df = next_df.cache()
current_df.count() # trigger materialization
result = current_df
Без checkpoint после 50 итераций логический план будет содержать 50 вложенных вызовов compute_next_iteration. Catalyst попытается оптимизировать этот огромный план - и либо потратит минуты, либо упадёт с StackOverflowError.
С checkpoint каждые 5 итераций максимальная глубина плана никогда не превысит 5 уровней. Plan planning работает быстро, пересчёт при сбое ограничен 5 шагами.
Мониторинг Actions в Spark UI¶
Spark UI предоставляет детальную информацию о выполнении Actions. Ключевые вкладки:
Вкладка SQL / DataFrame¶
Каждый Action, запущенный через DataFrame API, отображается как SQL-план. Нажав на конкретный план, вы видите:
- Physical Plan: дерево физических операторов с метриками. После
checkpoint()план начинается с оператораScan(чтение с диска) - lineage разорвана визуально - Metrics на каждом узле: rows processed, bytes read, spill size. Если видите большое spill - executor не хватает памяти, данные уходят на диск
Вкладка Stages¶
Каждая Stage показывает:
- Input Size / Records: сколько данных прочитано
- Shuffle Write / Read Size: сколько данных записано/прочитано при shuffle
- Task Duration: распределение времени выполнения Tasks. Если одна Task значительно дольше других - это признак data skew
- GC Time: время на garbage collection. Высокий GC (>10% от Task Duration) - признак memory pressure
foreachPartition в Spark UI¶
При выполнении foreachPartition Spark UI показывает Tasks в соответствующей Stage. Время Task включает:
- Task Deserialization Time: время десериализации функции
write_partitionиз байтов (должно быть мало) - Executor Run Time: собственно время выполнения функции (открытие соединения, batch write)
- Result Serialization Time: время сериализации результата для отправки на Driver (для foreachPartition обычно мало, возвращается только статус)
Если Executor Run Time равномерно распределён между Tasks - это признак хорошего распределения нагрузки. Если некоторые Tasks занимают в 10× больше - возможно, соединение к БД нестабильно или часть партиций значительно больше других.
Anti-patterns при использовании Actions¶
Антипаттерн 1: Actions в цикле без кэша¶
# ПЛОХО: каждый count() запускает полный DAG заново
for region in regions:
count = df.filter(f"region = '{region}'").count() # N полных прогонов
print(f"{region}: {count}")
# ХОРОШО: один Action с агрегацией
region_counts = df.groupBy("region").count().collect()
for row in region_counts:
print(f"{row['region']}: {row['count']}")
Антипаттерн 2: Collect на большом датасете¶
# ПЛОХО: стягиваем весь датасет на Driver для обработки
all_orders = df.collect()
for order in all_orders:
if order["amount"] > 1000:
process_large_order(order)
# ХОРОШО: фильтрация и обработка на Executors
df.filter("amount > 1000").foreachPartition(process_large_orders_partition)
Антипаттерн 3: Многократный collect одного DataFrame¶
# ПЛОХО: три collect = три полных прогона DAG
total = df.count()
samples = df.take(10)
result = df.collect()
# ХОРОШО: кэшируем перед серией Actions
df.cache()
total = df.count()
samples = df.take(10)
result = df.collect()
df.unpersist()
Антипаттерн 4: toPandas() без limit на большом датасете¶
# ПЛОХО: выгружаем весь датасет для графика
import matplotlib.pyplot as plt
df_pandas = df.toPandas() # потенциальный OOM
plt.scatter(df_pandas["x"], df_pandas["y"])
# ХОРОШО: sample + Arrow
spark.conf.set("spark.sql.execution.arrow.pyspark.enabled", "true")
df_sample = df.sample(fraction=0.01, seed=42).toPandas()
plt.scatter(df_sample["x"], df_sample["y"])
Антипаттерн 5: Длинный DAG без checkpoint¶
# ПЛОХО: граф растёт бесконтрольно
df = source_df
for company_id in company_ids: # 200 итераций
df = df.union(load_company_data(company_id)) # каждый union добавляет узел
result = df.groupBy("category").sum("revenue") # план из 200+ узлов
# ХОРОШО: checkpoint каждые N итераций
spark.sparkContext.setCheckpointDir("s3a://bucket/ckpt/")
df = source_df
for i, company_id in enumerate(company_ids):
df = df.union(load_company_data(company_id))
if (i + 1) % 20 == 0:
df = df.checkpoint() # сбрасываем план, максимальная глубина = 20
result = df.groupBy("category").sum("revenue")
Практика: рефакторинг опасного legacy кода¶
Разберём реальный legacy-скрипт с тремя типичными проблемами.
Проблема 1: Foreach вместо foreachPartition¶
Было (антипаттерн):
# Legacy-код: foreach с новым соединением на каждую строку
import psycopg2
def write_metric(row):
conn = psycopg2.connect(CONN_STRING)
conn.cursor().execute(
"INSERT INTO metrics (ts, name, value) VALUES (%s, %s, %s)",
(row["ts"], row["name"], row["value"])
)
conn.commit()
conn.close()
metrics_df.foreach(write_metric)
Стало (правильный паттерн):
def write_metrics_partition(rows):
"""Batch-запись партиции: одно соединение, N/batch_size запросов."""
conn = psycopg2.connect(CONN_STRING)
cursor = conn.cursor()
batch = []
try:
for row in rows:
batch.append((row["ts"], row["name"], row["value"]))
if len(batch) >= 500:
cursor.executemany(
"INSERT INTO metrics (ts, name, value) VALUES (%s, %s, %s)"
" ON CONFLICT (ts, name) DO UPDATE SET value = EXCLUDED.value",
batch
)
conn.commit()
batch = []
if batch:
cursor.executemany(
"INSERT INTO metrics (ts, name, value) VALUES (%s, %s, %s)"
" ON CONFLICT (ts, name) DO UPDATE SET value = EXCLUDED.value",
batch
)
conn.commit()
except Exception:
conn.rollback()
raise
finally:
cursor.close()
conn.close()
metrics_df.foreachPartition(write_metrics_partition)
Проблема 2: Итеративный алгоритм без checkpoint¶
Было (падает с StackOverflowError через ~80 итераций):
# Рекурсивная обработка иерархий организаций
employees = spark.table("hr.employees") # id, manager_id, level, salary
for level in range(100): # до 100 уровней иерархии
employees = employees.join(
employees.selectExpr("id as manager_id", "level as manager_level"),
on="manager_id",
how="left"
).withColumn("level", when(col("manager_level").isNotNull(),
col("manager_level") + 1).otherwise(0))
# После 80-100 итераций: java.lang.StackOverflowError
Стало (с checkpoint каждые 10 итераций):
from pyspark.sql.functions import when, col
spark.sparkContext.setCheckpointDir("s3a://data-lake-tmp/hr-hierarchy-ckpt/")
CKPT_INTERVAL = 10
employees = spark.table("hr.employees").cache()
for level in range(100):
employees_next = employees.join(
employees.selectExpr("id as manager_id", "level as manager_level"),
on="manager_id",
how="left"
).withColumn("level", when(col("manager_level").isNotNull(),
col("manager_level") + 1).otherwise(0))
if (level + 1) % CKPT_INTERVAL == 0:
# Разрываем lineage - план не вырастет больше 10 уровней
employees_next = employees_next.checkpoint()
employees.unpersist()
print(f"Checkpoint at level {level + 1}")
employees = employees_next.cache()
employees.count() # eager materialization
final_hierarchy = employees
Проблема 3: Небезопасный toPandas()¶
Было (периодически падает с OOM на Driver):
# Notebook для построения графиков продаж
import matplotlib.pyplot as plt
import seaborn as sns
sales_df = spark.table("analytics.daily_sales") # 50M строк, ~4 ГБ
df_pandas = sales_df.toPandas() # ← OOM при driver.memory=4g
sns.lineplot(data=df_pandas, x="date", y="revenue", hue="region")
plt.show()
Стало (Arrow + предагрегация + limit):
import matplotlib.pyplot as plt
import seaborn as sns
# Включить Arrow-оптимизацию
spark.conf.set("spark.sql.execution.arrow.pyspark.enabled", "true")
# Агрегировать ПЕРЕД выгрузкой - не тащить 50M строк ради линейного графика
daily_totals = (
spark.table("analytics.daily_sales")
.groupBy("date", "region")
.agg({"revenue": "sum", "orders": "count"})
.orderBy("date", "region")
# Ограничиваем: 365 дней × 10 регионов = 3650 строк максимум
.limit(5000)
)
# Проверяем размер перед выгрузкой (опционально, для надёжности)
row_count = daily_totals.count()
print(f"Выгружаем {row_count} строк в pandas")
df_pandas = daily_totals.toPandas() # Arrow: быстро и безопасно
sns.lineplot(data=df_pandas, x="date", y="sum(revenue)", hue="region")
plt.tight_layout()
plt.show()
Best Practices для production Spark Actions¶
Вызывайте Action только тогда, когда это необходимо. Каждый Action запускает полный DAG. Сгруппируйте несколько вычислений в один Action через агрегацию - вместо трёх count() используйте один .agg({"col_a": "count", "col_b": "sum"}).
Кэшируйте перед серией Actions. Если планируете вызвать несколько Actions на одном DataFrame - сначала cache() или persist(), потом Actions, потом unpersist(). Без кэша каждый Action читает данные заново.
Никогда не передавайте несериализуемые объекты через замыкание. Соединения к базам данных, файловые дескрипторы, ML-модели с нативными библиотеками (torch.nn.Module с CUDA) - всё это может быть несериализуемым. Создавайте такие объекты внутри foreachPartition.
Если df.explain() занимает больше экрана - делайте checkpoint. Длинный физический план - верный признак слишком длинного DAG. checkpoint() разрывает его и ускоряет как планирование, так и восстановление при сбоях.
Driver - координатор, не аналитическая система. Не стягивайте гигабайты на Driver через collect() или toPandas(). Driver должен управлять Job-ами и получать небольшие результаты агрегаций. Тяжёлые вычисления должны оставаться на Executor-ах.
Для toPandas() всегда используйте Arrow + limit/sample. Включите spark.sql.execution.arrow.pyspark.enabled = true один раз в начале session. Перед выгрузкой всегда применяйте limit() или убедитесь, что данные заведомо малы.
foreachPartition вместо foreach для внешних систем. Batch-запись с пулом соединений на уровне партиции - стандарт для любого external sink: PostgreSQL, ClickHouse, Elasticsearch, Kafka, REST API.
Освобождайте кэш явно. df.unpersist() после того, как DataFrame больше не нужен. Без явного освобождения кэш живёт до конца SparkSession и занимает память executor-ов, которая могла бы быть использована для других вычислений.