Классическая проблема RNN — их строго последовательная природа: каждый шаг зависит от предыдущего, из-за чего обучение и инференс плохо параллелятся и проигрывают трансформерам и SSM (например, Mamba). Но в SSM параллелизма добиваются ценой линейности рекуррентного перехода, ограничивая выразительность моделей.
Команда из Apple предлагает способ избежать этого компромисса: превратить применение RNN из итерационного процесса в решение системы нелинейных уравнений для всей последовательности.
Идея
Вместо того, чтобы последовательно пересчитывать каждое скрытое состояние через предыдущие, предлагают найти всё сразу.
Для решения системы используют два вложенных метода.
1. Внешний уровень — итерации метода Ньютона. На каждом шаге исходная система линеаризуется по якобианам нелинейной функции.
2. На внутреннем уровне — решение линейной системы, которое учитывает блочную би-диагональность матрицы в уравнении. Авторы замечают, что систему уравнений снова можно выразить рекуррентно. Но на этот раз каждый шаг рекурсии представлен в виде матричного умножения со сдвигом: Ax + b.
Рекуррентную систему такого вида можно решить алгоритмом parallel reduction за O(log₂(L)) шагов, где L — длина последовательности. Каждый шаг состоит из большого количества независимых задач, которые эффективно распаралливаются на GPU.
Таким образом, алгоритм хорошо загружает GPU вместо типичных «пары процентов утилизации» на длинных последовательностях.
Имплементация
К системной реализации авторы подошли максимально продакшн-ориентированно: сделали интеграцию с PyTorch + CUDA и полностью зафьюженные кернелы. Достаточно задать только рекуррентную формулу, остальное автоматизируется.
Сложность
На практике метод Ньютона быстро сходится — буквально за 3 итерации. Его результат эквивалентен обычному прогону RNN.
Итоговое время работы алгоритма можно оценить так:
latency = newton_iters ∙ log₂(L) ∙ (L / num_tasks_computed_in_parallel) ∙ time_per_task{}
Авторы репортят ускорение до космических x655 относительно наивного рекуррентного алгоритма.
Потенциальные проблемы
Дьявол кроется в последнем множителе оценки времени работы — time_per_task. В алгоритме parallel reduction любая отдельная подзадача подразумевает умножение двух матриц, каждая из которых либо якобиан нелинейной функции, либо результат перемножения якобианов.
В общем случае такая операция может быть довольно затратной и убивать выигрыш от параллелизации задачи. Авторы предпочли не упоминать об этом на постере в явном виде.
Именно поэтому в статье рассматривают RNN особого вида, где якобиан — либо диагональная, либо блочно-диагональная матрица с маленьким размером блока. Такие матрицы можно быстро умножать друг на друга.
Итого, применение метода оправдано только для тех RNN, чьи якобианы можно эффективно перемножать.
Разбор подготовил
#YaICLR26
Душный NLP