PyTorch 2.13: FlexAttention на Apple Silicon и fused ops для LLM
3328 коммитов от 526 контрибьюторов — таким получился PyTorch 2.13. Если 2.12 был про device-agnostic API и Microscaling, то 2.13 — релиз, который окончательно закрепляет трансформацию PyTorch из исследовательского фреймворка в production-платформу, одинаково удобную для обучения на кластере из тысяч GPU и инференса на MacBook Pro.
Три изменения выделяют этот релиз. FlexAttention наконец добрался до Apple Silicon с ускорением до 12 раз над стандартным SDPA для разреженных паттернов внимания. Fused-операция nn.LinearCrossEntropyLoss снижает пиковое потребление памяти в четыре раза при обучении языковых моделей с большим вокабуляром. А FSDP2 получил возможность перекрывать коммуникации all-gather и reduce-scatter через отдельную процесс-группу — это прямой буст пропускной способности распределённого обучения. Разбираем каждую фичу подробно и смотрим, что это значит для практики.
FlexAttention на Apple Silicon: 12x ускорения для sparse-паттернов
FlexAttention — это унифицированный API для кастомных паттернов внимания, где вместо написания CUDA-ядра вы описываете маску внимания как обычную Python-функцию, а компилятор генерирует fused kernel автоматически. До версии 2.13 это работало только на CUDA. Теперь FlexAttention доступен на Metal/MPS — то есть на всех Maках с чипами M1, M2, M3 и M4.
Реализация использует hand-written Metal-кernels как для sparse prefill, так и для decode-пути, включая поддержку Grouped Query Attention (GQA) и captured buffers. Результат на длинных разреженных паттернах впечатляет: при seq_len=32768 и sliding window=256 (плотность всего 0.8%) FlexAttention выполняется за ~35 мс против ~431 мс у стандартного SDPA — ускорение в 12.3 раза. На более коротких последовательностях (8192 токенов, окно 64) — всё равно солидные 4.15x.
Для разработчиков, которые запускают модели на Apple Silicon, это означает, что длинные документы с скользящим окном (RAG-системы, анализ кодовых баз, суммаризация больших текстов) теперь работают практически без штрафa за длину контекста. При этом для dense-паттернов SDPA по-прежнему быстрее — компилятор PyTorch достаточно умён, чтобы не навязывать FlexAttention там, где он не нужен.
Отдельного упоминания заслуживает детерминированный backward-путь для CUDA-версии FlexAttention. Ранее атомарные операции при akumulации градиентов dQ делали невозможным воспроизведение результатов — две прогонки на одинаковых входных данных давали немного разные градиенты. Новый путь compute_dq_write_order заменяет атомарики на предвычисленный порядок записи, гарантируя побитовое воспроизведение. Накладные расходы — менее 0.2% при S=32768, то есть детерминизм по сути бесплатный. Включается стандартным torch.use_deterministic_algorithms(True), без изменений в коде.
nn.LinearCrossEntropyLoss: fused операция для LLM-трейнинга
При обучении языковых моделей с большим вокабуляром (100K+ токенов) стандартный пайплайн выглядит так: финальный линейный слой проецирует hidden states на весь вокабуляр, формируя логиты размером [batch × seq_len × vocab_size], а затем CrossEntropyLoss вычисляет потери на этой матрице. Проблема в том, что матрица логитов может занимать десятки гигабайт GPU-памяти — особенно при больших батчах и длинных последовательностях.
nn.LinearCrossEntropyLoss решает это радикально: модуль объединяет линейную проекцию и вычисление cross-entropy в одну fused-операцию, которая обрабатывает вокабулярное измерение чанками, никогда не материализуя полную матрицу логитов. Результат — снижение пиковой памяти до 4 раз при числовой эквивалентности с нефьюздным путём.
Это drop-in замена: вместо последовательности nn.Linear(vocab_size) и nn.CrossEntropyLoss() вы используете один модуль, который принимает hidden states и целевые токены напрямую. Поддержка label smoothing, weight tying и z-loss regularization уже встроена, а интеграция с torch.compile даёт дополнительные оптимизации. Для команд, которые обучают LLM с нуля или дообучают на своих данных, это означает возможность увеличить batch size или длину последовательности без покупки дополнительного железа.
Практический пример: при обучении модели с вокабуляром 128K токенов, batch size 32 и seq_len 4096 — матрица логитов занимает ~62 ГБ в bf16. С fused-версией пиковая память снижается до ~15-16 ГБ, освобождая место для активаций, градиентов и optimizer states. Это особенно ценно при использовании FSDP или DeepSpeed, где перераспределение памяти между узлами — постоянная головная боль.
FSDP2: перекрывание коммуникаций для распределённого обучения
Fully Sharded Data Parallel (FSDP) — стандартный подход к обучению больших моделей на нескольких GPU. В каждой итерации FSDP выполняет all-gather (сбор параметров со всех узлов) перед forward pass, затем reduce-scatter (агрегация градиентов) после backward. Эти коммуникации блокируют вычисления — GPU простаивает, ожидая данные от сети.
FSDP2 в PyTorch 2.13 добавляет возможность перекрывать reduce-scatter и all-gather через выделенную процесс-группу. Вместо последовательного выполнения «forward → all-gather → compute → backward → reduce-scatter», система начинает следующую all-gather операцию параллельно с текущим reduce-scatter. Это классическая техника изMegatron-LM и DeepSpeed ZeRO, но теперь она встроена в ядро PyTorch как opt-in возможность.
Для кластеров с NVLink или InfiniBand это даёт измеримый прирост пропускной способности — чем дороже коммуникация относительно вычислений (то есть чем больше модель и меньше локальная память), тем заметнее эффект. При этом API остаётся совместимым: достаточно указать отдельную процесс-группу для overlap при инициализации FSDP2.
CuTeDSL: второй high-performance backend для Inductor
TorchInductor — компилятор PyTorch, который генерирует оптимизированные ядра для различных hardware-бэкендов. До 2.13 основным GPU-бэкендом был Triton. Теперь добавлен CuTeDSL — «нативный DSL», который генерирует kernels уровня CUTLASS (библиотека высокопроизводительных GEMM-операций от NVIDIA) напрямую, без зависимости от Triton.
CuTeDSL даёт более быструю компиляцию и в некоторых случаях более эффективные ядра для key GPU-операций — особенно для GEMM (обобщённое умножение матриц), которое составляет основу transformer-архитектур. Наличие двух независимых backend'ов также снижает рискVendor lock-in: если Triton деградирует на определённой архитектуре, CuTeDSL служит альтернативой.
Для большинства пользователей это прозрачное улучшение — torch.compile автоматически выбирает лучший backend. Но для команд, которые тюнят производительность критичных моделей, возможность переключаться между Triton и CuTeDSL даёт дополнительный рычаг оптимизации.
torchcomms: новый коммуникационный backend для distributed training
Крупномасштабное обучение (сотни и тысячи GPU) сталкивается с проблемами отказоустойчивости и отладки — если один узел падает, весь джоб останавливается, а поиск причины требует сложной диагностики. torchcomms — новый коммуникационный backend, разработанный специально для этих сценариев.
Он улучшает три аспекта: fault tolerance (быстрое восстановление после сбоя отдельных узлов без полной остановки), масштабируемость (оптимизированные алгоритмы для тысяч участников) и debuggability (встроенные инструменты для трассировки коммуникационных паттернов). Это не замена NCCL для стандартных сценариев, а дополнение для экстремальных масштабов, где каждая минута простоя стоит тысяч долларов.
Нативная загрузка Safetensors и поддержка Python 3.15
Safetensors стал де-факто стандартом распространения весов моделей — Hugging Face, Stability AI и другие используют его благодаря memory-mapped загрузке и отсутствию риска выполнения произвольного кода (в отличие от pickle). Теперь torch.load("model.safetensors") работает нативно — PyTorch определяет формат по расширению и возвращает тензоры напрямую, без необходимости устанавливать отдельную библиотеку safetensors.
Поддержка Python 3.15 (включая экспериментальный free-threaded билд 3.15t) пока ограничена Linux и самим torch (torchvision ещё не собран для 3.15), но это важный шаг к устранению GIL для параллельных workload'ов. Финальный релиз Python 3.15 запланирован на октябрь 2026 — PyTorch готовится заранее.
Расширенная поддержка платформ: ROCm, Arm, Intel XPU
PyTorch 2.13 продолжает движение к hardware-agnostic платформе. ROCm (AMD GPU) получает AOTriton 0.12b с нативной HIP CMake-интеграцией — это упрощает сборку и развёртывание на AMD-железе. Arm-платформа добавляет таргетинг torch.compile для Armv9-A, открывая оптимизации компилятора для edge-устройств. Intel XPU расширяет телеметрию — новые device-диагностические API помогают мониторить состояние accelerator'ов в production.
Отдельно стоит отметить интеграцию ExecuTorch в PyTorch Core. On-device инференс теперь не сторонний проект, а first-class capability основного фреймворка. Модель, обученная в PyTorch, может быть экспортирована на мобильные устройства и микроконтроллеры без отдельного pipeline'а конвертации.
Что это значит для практики
PyTorch 2.13 — не просто набор улучшений производительности. Это сигнализирование о зрелости: фреймворк больше не выбирает между «удобно для исследователей» и «эффективно для production». FlexAttention на Apple Silicon означает, что разработчик на MacBook может запускать модели с длинным контекстом без специализированного железа. Fused LinearCrossEntropyLoss означает, что обучение LLM на одном A100 теперь вмещает большие батчи. FSDP2 overlap означает, что распределённое обучение на кластере становится ближе к теоретическому пределу эффективности.
Если вы работаете с LLM — обновляйтесь. Fused ops и FlexAttention дают измеримую выгоду без изменений в архитектуре модели. Если вы работаете с distributed training — оцените torchcomms и FSDP2 overlap, особенно если масштаб превышает несколько сотен GPU. Если вы разрабатываете для edge — ExecuTorch в ядре PyTorch упрощает экспорт моделей на мобильные устройства.
Часто задаваемые вопросы
Нужно ли переписывать код для использования FlexAttention на Apple Silicon?
Нет, если вы уже используете стандартный attention API. FlexAttention работает через torch.nn.functional.scaled_dot_product_attention с кастомной маской, определённой как Python-функция. Для sliding window attention достаточно описать маску — компилятор сгенерирует Metal-kerne автоматически. Существующий код с SDPA продолжит работать без изменений, просто без ускорения sparse-паттернов.
Как включить детерминированный backward для FlexAttention на CUDA?
Достаточно вызвать torch.use_deterministic_algorithms(True) в начале скрипта. Это глобальная настройка, которая переключает все детерминированные реализации в PyTorch. Для FlexAttention это активирует путь compute_dq_write_order — предвычисленный порядок записи градиентов dQ вместо атомарных операций. Накладные расходы менее 0.2% при типичных длинах последовательностей.
CuTeDSL и Triton — это конкуренты или дополнения?
Это два независимых backend'а для TorchInductor, каждый со своими сильными сторонами. CuTeDSL генерирует kernels уровня CUTLASS с быстрой компиляцией, Triton предлагает более гибкий DSL для сложных паттернов. torch.compile автоматически выбирает backend на основе операции и hardware. Наличие двух путей снижает риск деградации производительности на определённых архитектурах — если один backend показывает худшие результаты, система может переключиться на другой.
nn.LinearCrossEntropyLoss совместим с FSDP и DeepSpeed?
Да, модуль разработан для работы в distributed-среде. Он интегрируется с torch.compile и поддерживает стандартные паттерны FSDP. Для DeepSpeed совместимость зависит от версии — рекомендуется проверить release notes DeepSpeed на предмет поддержки fused ops. В большинстве случаев это drop-in замена без изменений в конфигурации distributed training.
Итог
PyTorch 2.13 — релиз, который закрывает разрыв между исследовательской гибкостью и production-требованиями. FlexAttention на Apple Silicon делает длинный контекст доступным на потребительском железе, fused LinearCrossEntropyLoss режет память для LLM-трейнинга, а FSDP2 overlap выжимает максимум из распределённых кластеров. Плюс нативная загрузка safetensors, поддержка Python 3.15 и интеграция ExecuTorch — всё это делает PyTorch по-настоящему hardware-agnostic платформой. Обновляйтесь, тестируйте новые фичи и делитесь результатами на PyTorch Forums.