Narzut GPU: dlaczego nasz model MNIST trenuje się szybciej na CPU
GPU może przyspieszyć trening sieci neuronowej, wykonując wiele obliczeń równolegle. Mały model może jednak dostarczać zbyt mało pracy, aby zrekompensować koszty zlecania operacji, przesyłania danych i synchronizacji wyników.
Porównaliśmy nasz model MNIST na GPU Quadro RTX 5000 i procesorze Intel Core i9-10885H. W jednym przebiegu Kerasa CPU był szybszy: 4,4 s wobec 6,5 s na epokę. Osobny eksperyment dał zaledwie 0,14 s dla pętli treningowej JAX na CPU, ale obejmował inny zakres pracy niż Keras. Przyjrzymy się źródłom narzutu i temu, jakie wnioski można wyciągnąć z tych pomiarów.
Zacznijmy od tego, co dzieje się podczas treningu.
Rozkładamy jeden krok treningu
Krok treningu obejmuje przejście w przód, obliczenie gradientu straty, propagację wsteczną i aktualizację wag. Ten fragment oblicza bezpośrednio gradient softmaxu z entropią krzyżową, bez wyznaczania skalarnej wartości straty:
# Przejście w przód
z1 = xb @ w1 + b1 # mnożenie macierzy + bias (warstwa ukryta)
a1 = np.maximum(0, z1) # aktywacja ReLU
z2 = a1 @ w2 + b2 # mnożenie macierzy + bias (warstwa wyjściowa)
exp_z = np.exp(z2 - z2.max(axis=1, keepdims=True))
probs = exp_z / exp_z.sum(axis=1, keepdims=True) # softmax → prawdopodobieństwa
# Gradient straty
dz2 = probs.copy()
dz2[np.arange(bs), yb] -= 1 # jak bardzo przewidywania się pomyliły
dz2 /= bs
# Przejście w tył — obliczamy gradienty
dw2 = a1.T @ dz2 # gradient dla W2
da1 = dz2 @ w2.T # gradient płynący wstecz
dz1 = da1 * (z1 > 0) # gradient ReLU
dw1 = xb.T @ dz1 # gradient dla W1
# Aktualizujemy wagi
w1 -= lr * dw1
b1 -= lr * dz1.sum(axis=0)
w2 -= lr * dw2
b2 -= lr * dz2.sum(axis=0)Gdzie ucieka czas? Żeby to sprawdzić, owinęliśmy każdą operację w 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
# ... i tak dalej dla każdej operacjiOperacje GPU są asynchroniczne, dlatego dodaliśmy
tf.test.experimental.sync_devices(), aby poczekać na ich zakończenie. Synchronizacja po każdej operacji sama zwiększa narzut i uniemożliwia nakładanie się pracy, więc te czasy nie odpowiadają kosztowi zoptymalizowanej pętli treningowej.
Oto wyniki dla jednego kroku treningu (batch_size=32, średnia z 1000 kroków, czas w mikrosekundach μs — milionowych częściach sekundy, mniej znaczy lepiej):
| Operacja | CPU (NumPy) | GPU (TF) |
|---|---|---|
| fwd: X @ W1 + b1 (obrazy × wagi warstwy ukrytej + bias) | 162 μs | 541 μs |
| fwd: ReLU (zerowanie wartości ujemnych) | 9 μs | 173 μs |
| fwd: a1 @ W2 + b2 (warstwa ukryta × wagi wyjścia + bias) | 19 μs | 478 μs |
| fwd: softmax + strata (prawdopodobieństwa + błąd) | 30 μs | 780 μs |
| bwd: gradienty (ile każda waga wniosła do błędu) | 188 μs | 1271 μs |
| update: W -= lr*dW (korekta wag, by zmniejszyć błąd) | 345 μs | 1204 μs |
| RAZEM | 751 μs | 4447 μs |
W tym pomiarze pojedynczych operacji CPU jest szybszy. NumPy korzysta ze zoptymalizowanych procedur BLAS do mnożenia macierzy, a pomiary TensorFlow na GPU obejmują też zlecanie pracy i synchronizację. Instrukcje SIMD pozwalają CPU przetwarzać kilka wartości jedną instrukcją, a FMA łączy mnożenie z dodawaniem. Porównujemy więc implementacje wraz z metodą pomiaru, a nie samą szybkość arytmetyki obu urządzeń.
Spójrzmy na xb @ w1 + b1 i a1 @ w2 + b2 — mnożenia macierzy z każdej warstwy — i na to, dlaczego GPU są projektowane, by robić je szybko. Oto przejście w przód z naszej implementacji w NumPy:
class HiddenLayer:
def forward(self, x):
self.z = self.W @ x + self.b # mnożenie macierzy + bias
self.out = np.maximum(0, self.z) # aktywacja ReLU
return self.out
class OutputLayer:
def forward(self, x):
self.z = self.W @ x + self.b # mnożenie macierzy + bias
exp = np.exp(self.z - np.max(self.z))
self.probs = exp / np.sum(exp) # softmax → prawdopodobieństwa
return self.probsDla jednego obrazu self.W @ x mnoży macierz wag 128×784 przez wektor 784 pikseli. Każdy neuron oblicza iloczyn skalarny: 784 mnożenia i 783 dodawania. Po dodaniu biasu otrzymujemy 128 wartości przed aktywacją, a ReLU przekształca je w aktywacje warstwy ukrytej:
W (128 × 784) @ x (784,) + b (128,) → z (128,)Dla batcha 32 obrazów macierz X ma kształt (32, 784). Transpozycja wag przechowywanych w naszej implementacji dla pojedynczego obrazu daje macierz W1 o kształcie (784, 128). Iloczyn ma kształt (32, 128); następnie dodajemy bias i stosujemy ReLU:
X (32 × 784) W1 (784 × 128) result (32 × 128)
┌─────────────────┐ ┌──────────────────┐ ┌──────────────────┐
│ img1: p1 … p784│ │ n1 n2 … n128 │ │ img1: z1 … z128 │
│ img2: p1 … p784│ @ │ w w … w │ = │ img2: z1 … z128 │
│ ... │ │ ... ... ... │ │ ... │
│ img32: p1 … p784│ │ w w … w │ │ img32:z1 … z128 │
└─────────────────┘ └──────────────────┘ └──────────────────┘Przetwarzanie batcha wykorzystuje te same wagi dla 32 obrazów i zapisuje całą pracę jako jedno mnożenie macierzy. Ogranicza to liczbę osobnych wywołań oraz daje implementacji więcej możliwości ponownego użycia danych i równoległych obliczeń.
Wartości 4096 wyjść można obliczać niezależnie, ale zoptymalizowane mnożenie macierzy na GPU zwykle rozdziela fragmenty macierzy między grupy wątków. Z samej liczby wyjść nie można wywnioskować liczby aktywnych rdzeni ani multiprocesorów SM.
Dla orientacyjnego oszacowania przyjmijmy 3072 rdzenie FP32 pracujące z częstotliwością 1,8 GHz i dwie operacje zmiennoprzecinkowe na mnożenie z dodawaniem. Daje to teoretyczny szczyt około 11 TFLOPS. Rzeczywista wydajność zależy od wariantu GPU, częstotliwości i obciążenia.
Mnożenie macierzy wymaga około operacji zmiennoprzecinkowych, licząc mnożenie i dodawanie osobno. Przy 11 TFLOPS dolna granica czasu samej arytmetyki wynosi około 0,58 μs. Zakłada to szczytową wydajność i pomija dostęp do pamięci oraz zlecanie pracy. To nie jest zmierzony czas kernela. Podzielenie liczby operacji przez szczytową liczbę FLOPS daje czas, a nie procent wykorzystania GPU.
Na co zużywany jest czas
Krok treningu na GPU obejmuje kilka kosztów poza arytmetyką. Ich udział zależy od zadania i trybu wykonania:
1. Zlecanie kerneli. CPU i środowisko uruchomieniowe muszą przekazać pracę GPU. Przy małych operacjach może to zajmować znaczną część czasu. Jedna operacja frameworka może używać kilku kerneli, a kompilator może scalić kilka operacji w jeden kernel. Liczba wyrażeń w Pythonie nie wyznacza więc liczby uruchomień.
2. Przesyłanie danych i synchronizacja. Dane wejściowe mogą wymagać przesłania z pamięci CPU do GPU. Batch 32 obrazów po 784 piksele w formacie float32 zajmuje około 100 KB. Czas transferu zależy od połączenia, alokacji pamięci i synchronizacji. Podczas zwykłego treningu na GPU wagi i gradienty mogą pozostawać na urządzeniu; gradientów nie trzeba odsyłać do CPU po każdym kroku.
3. Praca Pythona i frameworka. Zlecanie operacji, alokacja pamięci, metryki i obsługa danych mają własny koszt. W trybie eager Python uczestniczy w poszczególnych operacjach. Wykonywanie grafu i kompilacja mogą ograniczyć tę powtarzaną pracę. XLA kompiluje funkcję, gdy jest potrzebna po raz pierwszy lub wymaga ponownego śledzenia, a nie przy każdej operacji w każdym kroku.
Profilowanie pomaga rozróżnić te koszty, o ile oddzielamy czas wywołań API na CPU, wykonanie na GPU i całkowity czas pomiaru.
Profilowanie narzutu CUDA
Pomiary pojedynczych operacji sugerują, że narzut ma znaczenie w tym zadaniu. Żeby go zbadać, potrzebujemy również profilu całego treningu.
Żeby to sprawdzić, użyliśmy profilera NVIDII nsys (Nsight Systems), który przechwytuje każde wywołanie API CUDA:
nsys profile -o keras-gpu-profile python mnist-keras.py
nsys stats --force-export=true keras-gpu-profile.nsys-repCUDA to warstwa programowa NVIDII między Twoim kodem a sprzętem GPU. Gdy TensorFlow chce pomnożyć macierze, nie rozmawia z GPU bezpośrednio — woła funkcje CUDA w rodzaju „zaalokuj pamięć”, „skopiuj te dane”, „uruchom ten kernel”. Każde z tych wywołań przechodzi przez sterownik i ma własny narzut. Profiler nsys rejestruje każde z nich, więc możemy dokładnie zobaczyć, gdzie ucieka czas.
Zmierzona epoka na GPU trwała 6,5 s. Poniżej znajduje się towarzyszące jej zestawienie wywołań API CUDA. Przy 48 000 obrazów treningowych i batchu 32 epoka ma 1500 kroków treningowych, ale zestawienie zawiera 1875 uruchomień grafu. Same liczby nie mówią, które wywołania należą do treningu, walidacji lub przygotowania.
| API CUDA | Funkcja | Wywołania | Suma czasów API |
|---|---|---|---|
cuCtxSetCurrent | ustawienie bieżącego kontekstu | 98 168 | 0,87 s |
cuEventRecord | zapis zdarzeń | 22 786 | 0,57 s |
cuMemcpyDtoHAsync | zlecenie kopiowania GPU→CPU | 6144 | 0,56 s |
cuMemcpyHtoDAsync | zlecenie kopiowania CPU→GPU | 5678 | 0,27 s |
cuGraphLaunch | uruchomienie grafu | 1875 | 0,16 s |
cuLaunchKernel | uruchomienie kernela | 1080 | 0,02 s |
| Suma podanych czasów API | 2,45 s |
To czasy wywołań API CUDA po stronie CPU, nie czasy kerneli GPU. Wywołania mogą nakładać się na pracę GPU i innych wątków CPU. Nie można odjąć ich sumy od 6,5 s i uznać reszty za narzut Pythona. Do takiego podziału potrzebna jest oś czasu z jasno określonym zakresem pomiaru. Nsight Systems rozróżnia czas API, oczekiwania w kolejce i wykonania kernela.
Wpisy cuGraphLaunch pokazują, że ten przebieg używał CUDA Graphs do zlecania zapisanej pracy. Odtwarzanie grafu może ograniczyć narzut, ale liczba wywołań nie dowodzi jednego uruchomienia na batch ani pomijalnego czasu obliczeń GPU.
We wcześniejszym porównaniu Kerasa epoka na CPU trwała 4,4 s, wobec 6,5 s na GPU. CPU unika wywołań CUDA i transferów do osobnego GPU, ale nadal ponosi koszty frameworka, planowania pracy i dostępu do pamięci.
Osobny mikrobenchmark TensorFlow mierzył (32, 784) @ (784, 128) bez synchronizacji po każdej operacji. Otrzymano następujące średnie czasy:
| Czas na jedno mnożenie macierzy | |
|---|---|
| GPU | 109 μs |
| CPU | 272 μs |
W tym teście wynik GPU jest około 2,5 raza szybszy. Oba czasy obejmują koszty oprogramowania i dostępu do pamięci, a nie tylko arytmetykę. Metoda pomiaru różni się też od tej z pierwszej tabeli, więc odwrócenie wyników nie jest sprzecznością.
Dołączony skrypt pomiaru skompilowanego kroku wykonuje różną pracę na obu urządzeniach: ścieżka GPU aktualizuje wagi, a ścieżka CPU tylko oblicza gradienty. Tych czasów nie można więc bezpośrednio porównywać. Rzetelny pomiar wymaga tych samych obliczeń i oczekiwania na zakończenie pracy obu urządzeń.
Podczas treningu nvidia-smi może pokazywać wykorzystanie GPU na poziomie 90–100%, nawet jeśli dane zadanie działa szybciej na CPU.
Wykorzystanie GPU w nvidia-smi oznacza część okresu pomiarowego, w której działa co najmniej jeden kernel. Nie mierzy odsetka zajętych rdzeni ani osiągniętego udziału szczytowych FLOPS. Ciągłe wykonywanie małych kerneli może więc dawać wysoki odczyt bez pełnego wykorzystania mocy obliczeniowej GPU.
Jak przyspieszyć jeden przebieg treningu
Wypróbowaliśmy pięć zmian: większe batche, Kerasa na CPU, pętlę NumPy, kompilację JAX i bezpośredni port do CuPy. Każda wpływa na inne składniki kosztu wykonania.
1. Większy rozmiar batcha
Większe batche zmniejszają liczbę aktualizacji na epokę. Dla 48 000 obrazów batch 32 daje 1500 kroków. Batch 4096 daje 12 kroków, jeśli uwzględnimy ostatni niepełny batch. Pętle NumPy i JAX w naszym eksperymencie odrzucają tę resztę, przetwarzając 11 pełnych batchy, czyli 45 056 obrazów.
Większy batch daje też każdemu mnożeniu macierzy więcej pracy, co może pomóc wykorzystać równoległość GPU. Liczba aktywnych SM zależy od wybranego kernela i wymaga pomiaru; nie wynika bezpośrednio z rozmiaru batcha.
Nie można jednak zwiększać rozmiaru batcha bez opamiętania — ma on bezpośredni wpływ na dokładność modelu. Większe batche dają gładsze, ale rzadsze aktualizacje gradientu, co może prowadzić do gorszej generalizacji. Właściwy rozmiar batcha musisz wyeksperymentować dla swojego konkretnego modelu.
2. Keras na CPU — pominąć narzut CUDA
Wyłączenie GPU przez tf.config.set_visible_devices([], 'GPU') przed inicjalizacją urządzeń uruchamia TensorFlow na CPU. Usuwa to pracę CUDA ze ścieżki wykonania, ale nie oznacza skrócenia czasu o sumę czasów wywołań API CUDA.
3. Czysty NumPy — pominąć także narzut frameworka
Pętla NumPy pomija środowisko uruchomieniowe TensorFlow i korzysta z BLAS do mnożenia macierzy. Nadal ponosi koszty wywołań Pythona, alokacji tablic i dostępu do pamięci.
4. JIT w JAX — skompilować cały krok
jit w JAX kompiluje funkcję, dzięki czemu kolejne wywołania nie muszą zlecać każdej operacji przez Pythona. Kompilacja może scalać operacje i zmniejszać narzut, ale jedna skompilowana funkcja nie musi być jednym kernelem GPU. Czas kompilacji należy oddzielić od pomiarów po rozgrzewce.
5. CuPy — a co, jeśli po prostu przeniesiemy NumPy na GPU?
Przenieśliśmy też do CuPy pętlę NumPy przetwarzającą obrazy pojedynczo. Małe operacje wektorowe wymagały wtedy wielokrotnego zlecania pracy GPU. Ta wersja potrzebowała 443 sekund na pięć epok, ponad sześciokrotnie więcej niż odpowiadająca jej pętla NumPy. Wynik dotyczy tej implementacji; wersja CuPy przetwarzająca całe batche byłaby innym porównaniem.
Składamy to razem
Poniższa tabela przedstawia osobny eksperyment dla różnych rozmiarów batcha. Zakres pracy nie jest jednakowy: Keras uwzględnia walidację i metryki, a pętle NumPy i JAX mierzą tylko trening i pomijają niepełne batche. To czasy konkretnych skryptów, a nie kontrolowany ranking szybkości frameworków.
| Rozmiar batcha | Keras GPU | Keras CPU | Czysty NumPy | JAX (CPU, JIT) |
|---|---|---|---|---|
| 32 | 6,26 s | 11,31 s | 2,71 s | 2,37 s |
| 128 | 2,74 s | 2,21 s | 4,70 s | 1,03 s |
| 512 | 2,64 s | 1,80 s | 1,30 s | 0,64 s |
| 1024 | 2,47 s | 1,43 s | 0,87 s | 0,63 s |
| 2048 | 2,32 s | 1,42 s | 0,84 s | 0,48 s |
| 4096 | 2,32 s | 1,39 s | 1,03 s | 0,14 s |
Wynik CPU przy batchu 32 wynosi tu 11,31 s, a nie wcześniejsze 4,4 s. Nie należy łączyć tych osobnych pomiarów w jedną deklarację przyspieszenia. Kontrolowane porównanie wymaga tych samych danych, aktualizacji, walidacji, rozgrzewki i granic pomiaru.
Keras na GPU zbliża się w tym eksperymencie do 2,3 s. Sama tabela nie pokazuje, który składnik ogranicza dalszą poprawę.
Keras na CPU skraca czas z 11,31 s do 1,39 s wraz ze wzrostem batcha. Pomiar obejmuje całe wywołanie fit() użyte w eksperymencie.
NumPy osiąga 0,84 s przy batchu 2048 i 1,03 s przy 4096. Same czasy nie wskazują przyczyny spowolnienia.
JAX osiąga 0,14 s przy batchu 4096 w pętli obejmującej wyłącznie trening. Iloraz 6,26 s i 0,14 s wynosi około 45, ale łączy różne batche i zakresy pracy. Nie dowodzi 45-krotnego przyspieszenia równoważnego treningu ani osiągnięcia tej samej trafności walidacyjnej.
Dla tego małego modelu warto sprawdzić ograniczenie powtarzanych wywołań i zmianę batcha przed zmianą sprzętu. Porównuj czas potrzebny do osiągnięcia tej samej jakości walidacyjnej, a także czas epoki.
Zrównoleglanie przy przeszukiwaniu hiperparametrów
Powyższe podejścia przyspieszają jeden przebieg treningu. Ale gdy przeszukujesz hiperparametry — próbujesz 5 współczynników uczenia albo 4 architektury — każdy przebieg jest całkowicie niezależny. Nie dzielą wag, gradientów ani stanu.
vmap w JAX dodaje do funkcji aktualizacji wymiar modeli, pozwalając jednym wywołaniem aktualizować pięć niezależnych zestawów parametrów na tych samych obrazach. W połączeniu z jit pozwala uniknąć pętli Pythona po modelach. Kompilator wybiera sposób wykonania operacji; nie gwarantuje jednego kernela ani jednoczesnego wykonywania wszystkich modeli.
W tym przykładzie wszystkie modele mają te same kształty parametrów. Kod pokazuje zwektoryzowaną aktualizację i zakłada, że x_batch oraz y_batch są już przygotowane:
import jax
import jax.numpy as jnp
from jax import vmap, jit, grad, random
def init_params(key):
# Ten sam model co wcześniej: 784→128→10, losowe wagi
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):
# Przejście w przód + strata entropii krzyżowej — ta sama matematyka co w wersji NumPy
w1, b1, w2, b2 = params
h = jnp.maximum(0, x @ w1 + b1) # warstwa ukryta + ReLU
logits = h @ w2 + b2 # warstwa wyjściowa
log_probs = jax.nn.log_softmax(logits, axis=-1)
return -jnp.mean(log_probs[jnp.arange(y.shape[0]), y])
def sgd_step(params, x, y, lr):
# Jeden krok treningu: policz gradienty, zaktualizuj wagi
grads = grad(loss_fn)(params, x, y) # JAX automatycznie różnicuje loss_fn
return tuple(p - lr * g for p, g in zip(params, grads))
# 5 współczynników uczenia, 5 zestawów wag, trenowanych jednocześnie
lr_array = jnp.array([0.001, 0.01, 0.1, 1.0, 10.0])
keys = random.split(random.PRNGKey(42), len(lr_array))
batched_params = vmap(init_params)(keys)
batched_step = jit(vmap(sgd_step, in_axes=(0, None, None, 0)))
batched_params = batched_step(batched_params, x_batch, y_batch, lr_array)Inną możliwością jest trenowanie każdej konfiguracji w osobnym procesie. Ten szkic zakłada dostępność Kerasa i danych treningowych. Uruchom go jako skrypt z warunkiem if __name__ == "__main__", a dobierając liczbę procesów, uwzględnij pamięć i wątki CPU używane przez każdy z nich:
import multiprocessing as mp
def train_one_config(args):
# Trenuje jeden model z jednym współczynnikiem uczenia — działa we własnym procesie
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
if __name__ == "__main__":
with mp.get_context("spawn").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]])Co warto zapamiętać
- Mały model może działać szybciej na CPU. Pierwsze porównanie Kerasa dało 4,4 s na epokę na CPU i 6,5 s na GPU; osobny eksperyment przyniósł inne czasy.
- Szczytowe TFLOPS nie mierzą wykorzystania GPU. Dolna granica czasu arytmetyki nie mówi, ile trwa kernel ani ile rdzeni jest aktywnych.
- Mierz pełny zakres pracy. Synchronizacja, walidacja, pomijanie batchy i kompilacja wpływają na znaczenie wyniku.
- Porównuj równoważny trening. Większe batche i skompilowane pętle mogą skracać czas, ale przed ogłoszeniem przyspieszenia porównaj jakość walidacyjną.
- Sprawdź dostępne urządzenia. Sama liczba parametrów nie określa, czy szybszy będzie CPU, czy GPU.
Word2vec to kolejny przykład znaczenia sposobu dostępu do danych. Przy negative sampling aktualizacja obejmuje wybrane wiersze macierzy embeddingów zamiast dużych mnożeń gęstych macierzy. Względna szybkość CPU i GPU zależy od grupowania tych odczytów i aktualizacji; nie jest to uniwersalna przewaga CPU.