Накладные расходы GPU: почему наша модель MNIST обучается быстрее на CPU

Машинное обучение и GPU кажутся неразделимыми — каждый туториал, каждый облачный провайдер, каждое руководство «как начать» направляет вас к GPU-инстансам. Логика проста: данные поступают на вход, тысячи ядер обрабатывают их параллельно, обучение идёт быстро. Но здесь предполагается, что все эти ядра действительно используются, — а для небольших задач это не так.

Мы обучили нашу модель MNIST на локальной рабочей станции с GPU Quadro RTX 5000 и процессором Intel Core i9-10885H, а затем только на CPU. CPU оказался быстрее: 4,4 с против 6,5 с на эпоху. В этой статье объясняется, почему, и показывается, как мы довели время обучения до 0,14 с, сокращая накладные расходы вместо апгрейда железа.

Начнём с того, что происходит во время обучения.

Разбираем один шаг обучения

Каждый шаг обучения проходит четыре фазы: прямой проход, потери, обратный проход, обновление весов:

# Прямой проход
z1 = xb @ w1 + b1                                # умножение матриц + смещение (скрытый слой)
a1 = np.maximum(0, z1)                            # активация ReLU
z2 = a1 @ w2 + b2                                # умножение матриц + смещение (выходной слой)
exp_z = np.exp(z2 - z2.max(axis=1, keepdims=True))
probs = exp_z / exp_z.sum(axis=1, keepdims=True)  # softmax → вероятности

# Градиент потерь
dz2 = probs.copy()
dz2[np.arange(bs), yb] -= 1                       # насколько предсказания промахнулись
dz2 /= bs

# Обратный проход — вычисляем градиенты
dw2 = a1.T @ dz2                                  # градиент для W2
da1 = dz2 @ w2.T                                  # градиент, текущий назад
dz1 = da1 * (z1 > 0)                              # градиент ReLU
dw1 = xb.T @ dz1                                  # градиент для W1

# Обновляем веса
w1 -= lr * dw1
b1 -= lr * dz1.sum(axis=0)
w2 -= lr * dw2
b2 -= lr * dz2.sum(axis=0)

Куда уходит время? Чтобы выяснить это, мы обернули каждую операцию в time.perf_counter():

t = time.perf_counter()
z1 = xb @ w1 + b1
timings["fwd: X @ W1 + b1"] += time.perf_counter() - t

t = time.perf_counter()
a1 = np.maximum(0, z1)
timings["fwd: ReLU"] += time.perf_counter() - t
# ... и так далее для каждой операции

При запуске на GPU нам пришлось добавлять tf.test.experimental.sync_devices() между операциями — без этого операции GPU выстраиваются в очередь асинхронно, и замер фиксирует только отправку задания, а не фактическое выполнение.

Вот результаты для одного шага обучения (batch_size=32, среднее по 1000 шагам, время в микросекундах μs — миллионных долях секунды, меньше — лучше):

ОперацияCPU (NumPy)GPU (TF)
fwd: X @ W1 + b1 (изображения × веса скрытого слоя + смещение)162 μs541 μs
fwd: ReLU (обнуление отрицательных)9 μs173 μs
fwd: a1 @ W2 + b2 (скрытый слой × веса выхода + смещение)19 μs478 μs
fwd: softmax + потери (вероятности + ошибка)30 μs780 μs
bwd: градиенты (сколько каждый вес добавил к ошибке)188 μs1 271 μs
update: W -= lr*dW (корректируем веса, чтобы снизить ошибку)345 μs1 204 μs
ИТОГО751 μs4 447 μs

Первое, что бросается в глаза: CPU быстрее на каждой без исключения операции — 751 μs в сумме против 4 447 μs. Операция, которая должна была бы выиграть от распараллеливания на GPU больше всех, — это умножение матриц X @ W1: это тысячи независимых скалярных произведений, которые могли бы считаться на тысячах ядер одновременно, ровно та работа, под которую GPU и создавались. И всё же в таблице выше она по-прежнему медленнее на GPU: 162 μs на CPU против 541 μs на GPU. CPU здесь быстр потому, что NumPy обращается напрямую к BLAS (Basic Linear Algebra Subprograms) — сильно оптимизированным процедурам на C/Fortran, которые используют специфичные для CPU SIMD-инструкции (AVX2 обрабатывает 8 чисел с плавающей точкой за такт, FMA сливает умножение и сложение в одну инструкцию). GPU медленнее не потому, что его арифметика медленнее, а потому, что каждая операция платит накладные расходы CUDA (запуск ядра, синхронизация памяти, переключение контекста). Когда эти накладные расходы отнимают больше времени, чем сами вычисления, GPU в итоге оказывается медленнее, хотя считает быстрее.

Посмотрим на xb @ w1 + b1 и a1 @ w2 + b2 — умножения матриц из каждого слоя — и на то, почему GPU спроектированы делать их быстро. Вот прямой проход из нашей реализации на NumPy:

class HiddenLayer:
    def forward(self, x):
        self.z = self.W @ x + self.b        # умножение матриц + смещение
        self.out = np.maximum(0, self.z)     # активация ReLU
        return self.out

class OutputLayer:
    def forward(self, x):
        self.z = self.W @ x + self.b        # умножение матриц + смещение
        exp = np.exp(self.z - np.max(self.z))
        self.probs = exp / np.sum(exp)       # softmax → вероятности
        return self.probs

Каждый @ — это умножение матриц. Для одного изображения self.W @ x умножает матрицу весов на пиксели изображения. У каждого из 128 нейронов 784 веса — по одному на входной пиксель. Каждый нейрон вычисляет скалярное произведение своих 784 весов с 784 входными значениями (пикселями изображения) — это 784 умножения и 783 сложения на нейрон. Результат — 128 активаций: они получаются от 128 нейронов слоя, каждый из которых выдаёт по одному значению из своего скалярного произведения. Итого 128 независимых скалярных произведений на одно изображение:

      x (784,)                    W (128 × 784)              result (128,)
 ┌─────────────┐           ┌──────────────────┐       ┌──────────────┐
 │ p1 p2 … p784│     @     │ n1:  w1 w2 … w784│   =   │ a1           │
 └─────────────┘           │ n2:  w1 w2 … w784│       │ a2           │
                           │ ...              │       │ ...          │
                           │ n128:w1 w2 … w784│       │ a128         │
                           └──────────────────┘       └──────────────┘

Слева — одно изображение из 784 пикселей, в середине — 128 нейронов, справа — 128 активаций.

Но фреймворки вроде Keras не обрабатывают изображения по одному — они складывают весь батч в одну матрицу и умножают всё сразу. При размере батча 32 X — это 32 строки × 784 столбца, и та же матрица весов даёт 32 × 128 = 4096 активаций за одну операцию:

      X (32 × 784)                W1 (784 × 128)            result (32 × 128)
 ┌─────────────────┐         ┌──────────────────┐       ┌──────────────────┐
 │ img1:  p1 … p784│         │  n1   n2  … n128 │       │ img1: a1  … a128 │
 │ img2:  p1 … p784│    @    │  w    w   …  w   │   =   │ img2: a1  … a128 │
 │ ...             │         │  ...  ...    ... │       │ ...              │
 │ img32: p1 … p784│         │  w    w   …  w   │       │ img32:a1  … a128 │
 └─────────────────┘         └──────────────────┘       └──────────────────┘

Слева — 32 изображения по 784 пикселя каждое, в середине — 784 веса на нейрон, справа — 32 × 128 активаций.

Веса общие — это как если бы у вас было 32 копии одной и той же модели, работающие на 32 изображениях одновременно, только всё делается одной операцией, а не 32 отдельными. Один запуск ядра, одно чтение матрицы весов из памяти — и все 4096 результатов вычисляются параллельно.

Теперь интересный вопрос: сколько из этих 4096 скалярных произведений GPU может вычислить одновременно? Поскольку каждое из них полностью независимо — активация нейрона 20 для изображения 5 не зависит ни от одного другого результата, — GPU может назначить каждому отдельный поток.

У RTX 5000 48 потоковых мультипроцессоров (SM), в каждом по 64 ядра CUDA — всего 3072 ядра. Каждое ядро работает на частоте ~1,8 ГГц и может выполнять умножение со сложением за такт, что даёт теоретический пик: 3072 ядра × 1,8 млрд тактов/с × 2 операции/такт = ~11 TFLOPS (триллионов операций с плавающей точкой в секунду).

Наше умножение матриц даёт 4096 скалярных произведений, каждое из которых требует 784 умножений со сложением: 4096 × 784 = ~3,2 миллиона операций суммарно. GPU способен выполнять 11 триллионов операций в секунду, а мы просим всего 3,2 миллиона — это 0,3 микросекунды реальных вычислений, то есть всего 0,00003% его мощности. GPU заканчивает арифметику за 0,3 μs, а затем ждёт ~540 μs следующего запуска ядра. Он работает 0,06% времени.

Куда на самом деле уходит время

Обучение нейросети на GPU — это не только умножение матриц. Каждый шаг обучения представляет собой последовательность операций, и для небольших моделей сама арифметика оказывается наименьшей её частью:

1. Накладные расходы на запуск ядра. Каждая операция GPU — умножение матриц, функция активации, вычисление потерь — это ядро (kernel), которое CPU должен спланировать и запустить на GPU. Каждый запуск несёт фиксированные накладные расходы порядка 5–15 микросекунд. Для большого умножения матриц, занимающего миллисекунды, это пренебрежимо мало. Для нашего крошечного умножения 32×784 × 784×128, которое завершается за микросекунды, расходы на запуск могут превысить сами вычисления. Один только наш прямой проход включает умножение матриц, добавление смещения, ReLU, ещё одно умножение матриц, ещё одно смещение и softmax — как минимум 6 запусков ядер, ещё до начала обратного распространения.

2. Задержка передачи памяти. Данные должны перемещаться из памяти CPU в память GPU (а градиенты — обратно). У этой передачи есть фиксированная задержка: время на настройку DMA-передачи, прохождение по шине PCIe и сигнал о завершении. При размере батча 32 и 784-мерных входах мы передаём около 100 КБ на батч. Шина PCIe способна двигать 32 ГБ/с, так что сама передача занимает ~3 микросекунды, — но накладные расходы на подготовку в 10–20 раз больше.

3. Накладные расходы Python и фреймворка. Keras/TensorFlow добавляют свой слой опосредования. Каждая операция проходит через Python, рантайм TensorFlow, компиляцию XLA (компилятор, оптимизирующий граф вычислений при первом запуске), выделение памяти и синхронизацию. Для больших операций эти расходы незаметны. Для маленьких они и есть узкое место.

Хорошая новость в том, что мы можем точно измерить, сколько времени отнимает каждый вид накладных расходов.

Профилируем накладные расходы CUDA

Наши замеры по операциям показали, что GPU тратит больше времени, и мы посчитали, что сама арифметика занимает всего 0,3 μs — то есть GPU должен большую часть времени простаивать. Так на что же именно уходят остальные 540 μs?

Чтобы выяснить это, мы воспользовались профилировщиком NVIDIA nsys (Nsight Systems), который перехватывает каждый вызов CUDA API:

nsys profile -o keras-gpu-profile python mnist-keras.py
nsys stats --force-export=true keras-gpu-profile.nsys-rep

CUDA — это программный слой NVIDIA между вашим кодом и оборудованием GPU. Когда TensorFlow хочет умножить матрицы, он не обращается к GPU напрямую — он вызывает функции CUDA вроде «выдели память», «скопируй эти данные», «запусти это ядро». Каждый такой вызов идёт через драйвер и несёт собственные накладные расходы. Профилировщик nsys записывает каждый из них, так что мы можем точно увидеть, куда уходит время.

На этот раз вместо замеров отдельных операций мы профилировали целую эпоху model.fit() — все 1500 шагов. Эпоха заняла 6,5 с на GPU. Вот куда ушло это время:

ОперацияТип накладных расходовВызовыВсего времени
cuCtxSetCurrent (переключение контекста)запуск ядра98 1680,87 с
cuEventRecord (события замера/синхронизации)запуск ядра22 7860,57 с
cuMemcpyDtoHAsync (копирование GPU→CPU)передача памяти6 1440,56 с
cuMemcpyHtoDAsync (копирование CPU→GPU)передача памяти5 6780,27 с
cuGraphLaunch (выполнение скомпилированного графа)запуск ядра1 8750,16 с
cuLaunchKernelзапуск ядра1 0800,02 с
Итого накладные расходы CUDA~2,7 с

Ни одна из этих строк CUDA не является собственно умножением матриц — все они относятся к управляющим накладным расходам вокруг него. Сами вычисления (умножения матриц, ReLU, softmax) происходят на GPU после того, как cuGraphLaunch отправляет работу, но для наших крошечных матриц это настолько быстро, что даже не попадает в значимые позиции.

Собираем всё вместе — вот куда ушла вся эпоха длиной 6,5 с:

Время
Накладные расходы CUDA (таблица выше)~2,7 с
Накладные расходы Python/фреймворка (не фиксируются nsys)~3,8 с
Реальная арифметика на GPUпренебрежимо
Всего за эпоху6,5 с

По сути все 6,5 с — это накладные расходы. Keras уже очень хорошо оптимизирован: он использует CUDA Graphs (cuGraphLaunch — 1875 вызовов, по одному на батч), чтобы заранее скомпилировать весь прямой и обратный проход и воспроизводить его без диспетчеризации по операциям. Но даже с этой оптимизацией GPU так и не получает шанса отбить накладные расходы за счёт более быстрых вычислений.

Для сравнения: Keras на CPU проходит ту же эпоху за 4,4 с — быстрее, чем 6,5 с на GPU, потому что накладных расходов CUDA нет вовсе. XLA компилирует в нативный код, и CPU выполняет арифметику прямо в своей памяти.

Важно подчеркнуть, что с самой арифметикой никаких сюрпризов нет — GPU существенно быстрее, как и ожидалось. Мы замерили ключевую операцию — одно умножение матриц (32, 784) @ (784, 128), один батч через первый слой:

Время на одно умножение матриц
GPU109 μs
CPU272 μs

GPU быстрее в 2,5 раза — но учтите, что даже эти 109 μs уже включают накладные расходы CUDA на запуск одной этой операции. Чистая арифметика заняла бы ~0,3 μs (как мы посчитали ранее из TFLOPS); остальные ~108 μs — накладные расходы на этот единственный запуск ядра. На CPU все 272 μs — это арифметика: никакого промежуточного слоя накладных расходов.

Полный шаг обучения включает несколько умножений матриц плюс активации, потери, градиенты и всё управление CUDA вокруг каждого из них. Когда мы замеряем скомпилированный train_step со всеми накладными расходами, разрыв исчезает полностью:

На шаг
GPU1,27 мс (~0,1 мс вычислений + ~1,2 мс накладных расходов)
CPU1,30 мс (всё вычисления, накладных расходов нет)

GPU считает быстрее, но остальное время тратит на накладные расходы, в итоге выходя примерно на ту же скорость на шаг. А в полном конвейере model.fit() — с загрузкой данных, метриками и колбэками — CPU и вовсе выигрывает: 4,4 с против 6,5 с на эпоху.

Любопытная деталь: если запустить nvidia-smi во время обучения, вы можете увидеть загрузку GPU на уровне 90–100%. Похоже, что GPU полностью занят, — так почему же он медленнее CPU?

Потому что nvidia-smi сообщает долю времени, когда на GPU выполняется хотя бы одно ядро, а не число активных ядер. Если крошечные ядра запускаются одно за другим без простоев, показывается ~100% загрузки, хотя подавляющая часть ядер в любой момент простаивает. Как мы видели выше, наша модель использует всего 0,00003% мощности GPU.

Как ускорить один запуск обучения

Теперь, когда мы понимаем природу накладных расходов, как их сократить? Есть несколько направлений атаки: уменьшить число шагов (больший размер батча), устранить накладные расходы CUDA (работать на CPU), устранить накладные расходы фреймворка (обойтись без TensorFlow), скомпилировать весь шаг в одну операцию (JIT в JAX) или наивно переехать на GPU (CuPy — спойлер: получается хуже). Мы попробовали все пять.

1. Больший размер батча

Если накладные расходы платятся за каждый шаг, очевидное решение — шагов должно быть меньше. Больший размер батча означает больше примеров на шаг, поэтому та же эпоха требует меньше шагов, а каждый шаг платит налог накладных расходов лишь один раз, независимо от размера батча. При 48 000 обучающих примеров (60 тыс. минус 20% на валидацию) и размере батча 32 одна эпоха — это 48 000 / 32 = 1500 шагов, где каждый шаг есть один полный прямой + обратный проход и обновление на одном батче. При размере батча 4096 это всего 11 шагов.

Большие батчи также дают GPU больше работы на шаг: как мы видели ранее, размер батча 32 держит активными лишь ~3 из 48 SM, тогда как размер батча 2048 занимает ~42. Так что большие батчи помогают дважды: меньше шагов, за которые платятся накладные расходы, и больше реально работающих ядер GPU на каждом шаге.

Однако увеличивать размер батча без разбора нельзя — он напрямую влияет на точность модели. Большие батчи дают более гладкие, но менее частые обновления градиента, что может ухудшить обобщение. Правильный размер батча нужно подбирать экспериментально под вашу конкретную модель.

2. Keras на CPU — обойти накладные расходы CUDA

Как мы видели, простое отключение GPU (tf.config.set_visible_devices([], 'GPU')) устраняет 2,7 с накладных расходов CUDA. XLA компилирует в нативный код, и CPU выполняет арифметику напрямую.

3. Чистый NumPy — обойти ещё и накладные расходы фреймворка

Пойдём дальше: можно вообще отказаться от TensorFlow. Цикл обучения на чистом NumPy обращается напрямую к процедурам BLAS, без диспетчеризации фреймворка между операциями.

4. JIT в JAX — скомпилировать весь шаг

Компилятор jit из JAX трассирует весь шаг обучения и компилирует его в одну оптимизированную нативную функцию. Вместо того чтобы Python отправлял каждую операцию по очереди, скомпилированная функция выполняет их все одним слитым вызовом с почти нулевыми накладными расходами.

5. CuPy — а что если просто перенести NumPy на GPU?

Мы попробовали и наивный подход: заменить import numpy as np на import cupy as cp и запустить тот же код на GPU. Поскольку наша реализация обрабатывает примеры по одному в Python-цикле, каждая крошечная операция (сложение вектора из 128 элементов, softmax из 10 элементов) превращается в отдельный запуск ядра GPU. Результат: 443 секунды на 5 эпох — более чем в 6 раз медленнее NumPy на CPU. Простой перенос кода на GPU без переосмысления схемы доступа делает хуже, а не лучше.

Собираем всё вместе

Мы протестировали подходы 1–4 на разных размерах батча (CuPy оказался слишком медленным, чтобы его включать). В каждой ячейке — время одной эпохи:

Размер батчаKeras GPUKeras CPUЧистый NumPyJAX (CPU, JIT)
326,26 с11,31 с2,71 с2,37 с
1282,74 с2,21 с4,70 с1,03 с
5122,64 с1,80 с1,30 с0,64 с
10242,47 с1,43 с0,87 с0,63 с
20482,32 с1,42 с0,84 с0,48 с
40962,32 с1,39 с1,03 с0,14 с

Чтение по строкам показывает эффект смены подхода (меньше накладных расходов на шаг). Чтение по столбцам показывает эффект больших батчей (меньше шагов на эпоху). Эти два эффекта складываются.

Keras на GPU упирается в ~2,3 с независимо от размера батча — у накладных расходов CUDA есть фиксированный пол, который большие батчи убрать не могут.

Keras на CPU идёт с 11,3 с вниз до 1,4 с — накладных расходов CUDA нет, но накладные расходы самого фреймворка TensorFlow задают свой пол.

Чистый NumPy достигает 0,84 с при размере батча 2048 — никакого фреймворка, только вызовы BLAS. На 4096 он снова замедляется из-за давления на память.

JAX доходит до 0,14 с при размере батча 4096 — в 45 раз быстрее, чем Keras на GPU. JIT-компиляция сливает весь шаг в один нативный вызов с почти нулевыми накладными расходами.

Урок таков: путь к более быстрому обучению — не более мощное железо, а меньше накладных расходов. Каждый подход в таблице снимает один слой накладных расходов, а большие батчи сокращают, сколько раз вы платите за то, что осталось.

Используем распараллеливание для перебора гиперпараметров

Подходы выше ускоряют один запуск обучения. Но когда вы перебираете гиперпараметры — пробуете 5 скоростей обучения или 4 архитектуры, — каждый запуск полностью независим. Они не делят между собой ни веса, ни градиенты, ни состояние.

vmap (векторизованный map) из JAX может обучать все 5 моделей одновременно за один прямой и обратный проход. Это не то же самое, что пакетная обработка в Keras, которая прогоняет много изображений через одну модель. vmap прогоняет одни и те же изображения через много моделей — у каждой свои веса и своя скорость обучения — одной слитой операцией. Под капотом XLA компилирует это в одно ядро: одно умножение матриц обслуживает все 5 прямых проходов, другое — все 5 обратных. Никакого Python-цикла, никаких накладных расходов на модель.

Плата за это: все модели должны иметь одинаковую архитектуру (одинаковые формы параметров), и цикл обучения придётся переписать на JAX. Вот полный код:

import jax
import jax.numpy as jnp
from jax import vmap, jit, grad, random

def init_params(key):
    # Та же модель, что и раньше: 784→128→10, случайные веса
    k1, k2 = random.split(key)
    w1 = random.normal(k1, (784, 128)) * jnp.sqrt(2.0 / 784)
    b1 = jnp.zeros(128)
    w2 = random.normal(k2, (128, 10)) * jnp.sqrt(2.0 / 128)
    b2 = jnp.zeros(10)
    return (w1, b1, w2, b2)

def loss_fn(params, x, y):
    # Прямой проход + перекрёстная энтропия — та же математика, что в версии на NumPy
    w1, b1, w2, b2 = params
    h = jnp.maximum(0, x @ w1 + b1)       # скрытый слой + ReLU
    logits = h @ w2 + b2                    # выходной слой
    log_probs = logits - jnp.log(jnp.sum(jnp.exp(logits), axis=-1, keepdims=True))
    return -jnp.mean(log_probs[jnp.arange(y.shape[0]), y])

def sgd_step(params, x, y, lr):
    # Один шаг обучения: считаем градиенты, обновляем веса
    grads = grad(loss_fn)(params, x, y)     # JAX автоматически дифференцирует loss_fn
    return tuple(p - lr * g for p, g in zip(params, grads))

# 5 скоростей обучения, 5 наборов весов, обучаются одновременно
lr_array = jnp.array([0.001, 0.01, 0.1, 1.0, 10.0])
batched_step = jit(vmap(sgd_step, in_axes=(0, None, None, 0)))
batched_params = batched_step(batched_params, x_batch, y_batch, lr_array)

Для более простых случаев тот же результат даёт multiprocessing.Pool из Python — 5 отдельных процессов, каждый обучает одну конфигурацию:

import multiprocessing as mp

def train_one_config(args):
    # Обучаем одну модель с одной скоростью обучения — работает в своём процессе
    lr, X_train, y_train = args
    model = keras.Sequential([
        keras.layers.Dense(128, activation="relu", input_shape=(784,)),
        keras.layers.Dense(10, activation="softmax"),
    ])
    model.compile(optimizer=keras.optimizers.SGD(learning_rate=lr),
                  loss="sparse_categorical_crossentropy")
    history = model.fit(X_train, y_train, epochs=10, batch_size=32,
                        validation_split=0.2, verbose=0)
    return lr, history.history

with mp.Pool(5) as pool:
    results = pool.map(train_one_config,
                       [(lr, X_train, y_train) for lr in [0.001, 0.01, 0.1, 1.0, 10.0]])

Что стоит запомнить

  1. Для маленьких моделей CPU быстрее GPU. Наша модель MNIST (784→128→10, ~101 тыс. параметров) обучалась за 4,4 с на эпоху на CPU против 6,5 с на GPU. Арифметика GPU быстрее (109 μs против 272 μs на одно умножение матриц), но накладные расходы CUDA — запуски ядер, копирования памяти, переключения контекста — добавляют 2,7 с на эпоху, которых у CPU просто нет.
  2. GPU почти не используется. Наше умножение матриц задействует 0,00003% мощности RTX 5000. При batch_size=32 активны лишь ~3 из 48 SM. GPU заканчивает арифметику за 0,3 μs, а затем ждёт ~540 μs следующего запуска ядра.
  3. Показатель загрузки в nvidia-smi вводит в заблуждение. Он сообщает долю времени, когда выполняется хоть какое-то ядро, а не занятость ядер процессора GPU. Наша модель показывает ~100% загрузки, используя 0,00003% вычислительной мощности.
  4. Меньше накладных расходов лучше, чем мощнее железо. Перейдя с Keras на GPU к JIT в JAX на CPU с batch_size=4096, мы прошли путь от 6,5 с до 0,14 с на эпоху — ускорение в 45 раз на той же машине, без всякого GPU.
  5. Не хватайтесь за GPU, пока ваша задача этого не требует. Для моделей с менее чем ~500 тыс. параметров и размерами батча до 256 быстрый CPU окажется и дешевле, и быстрее.

Тот же самый сюжет широко известен в продакшене для word2vec: каждый шаг обучения затрагивает лишь горстку строк двух матриц эмбеддингов и выполняет порядка 10⁴ умножений со сложением — примерно в тысячу раз меньше вычислений на пример, чем прямой проход MNIST выше. Word2vec — это задача, в которой доминирует доступ к памяти и случайный поиск строк, а не плотное умножение матриц; именно такая форма проигрывает на GPU и выигрывает на CPU. Де-факто стандартная библиотека для word2vec, Gensim, доводит это до предела: внутренний цикл на Cython, блокировочно-свободные CPU-потоки в духе Hogwild!, никакого GPU ни на одном этапе — и она по-прежнему обгоняет GPU-реализации по времени на стенных часах при обучении таблицы поиска (V, d) по корпусу, спустя более десяти лет после того, как CUDA начала доминировать в глубоком обучении. Статья про word2vec подробно объясняет, почему.