06 · Sieci · 5 min czytania · Interaktywne · aktualizacja
Czym jest automatyczne różniczkowanie i jak PyTorch liczy gradienty?
W skrócie
Automatyczne różniczkowanie liczy dokładne pochodne programu: rozkłada go na proste operacje i łączy ich pochodne regułą łańcuchową, bez wzorów i przybliżeń.
Co to jest
Automatyczne różniczkowanie (ang. automatic differentiation, AD, w PyTorch autograd) to technika obliczania dokładnych wartości pochodnych funkcji zapisanej jako program. Rozkłada obliczenie na elementarne operacje (dodawanie, mnożenie, sinus, wykładnik), dla których pochodne są znane, i łączy je regułą łańcuchową. Wynik jest dokładny z dokładnością do arytmetyki zmiennoprzecinkowej, a koszt jest porównywalny z kosztem samego obliczenia funkcji.
AD nie jest różniczkowaniem symbolicznym, takim jak w programach algebry komputerowej, które produkują wzór pochodnej i mogą go rozdmuchać do ogromnych rozmiarów. Nie jest też różniczkowaniem numerycznym, czyli przybliżeniem (f(x + h) − f(x)) / h, obarczonym błędem zależnym od h. AD liczy wartość pochodnej w konkretnym punkcie, krok po kroku, razem z wartością funkcji.
To dzięki AD współczesne uczenie głębokie jest praktyczne. Badacz pisze tylko przejście w przód, nawet z pętlami i instrukcjami warunkowymi, a biblioteka sama dostarcza gradient. Bez tego każda nowa architektura wymagałaby ręcznego wyprowadzania i testowania wzorów na gradient.
Mechanizm — dlaczego tak działa
Każdy program liczący liczby jest ciągiem elementarnych kroków. Dla f(x, y) = x·y + sin(x) są to: v₁ = x·y, v₂ = sin(x), f = v₁ + v₂. Pochodna każdego kroku jest trywialna, a reguła łańcuchowa mówi, jak je połączyć. AD różni się od ręcznego liczenia tylko tym, że robi to mechanicznie i dla programów o milionach kroków.
Tryb w przód (forward mode) niesie razem z każdą wartością jej pochodną względem jednego wybranego wejścia. Wygodnie myśleć o tym jak o liczbach dualnych: zamiast x liczymy parę (x, x'), a każda operacja ma regułę dla obu części, np. (a, a')·(b, b') = (a·b, a'·b + a·b'). Jedno przejście daje pochodne względem jednego wejścia, ale za to wszystkich wyjść. Tryb w przód opłaca się, gdy wejść jest mało, a wyjść dużo.
Tryb odwrotny (reverse mode) najpierw liczy funkcję i zapisuje graf obliczeń (tzw. taśmę), potem przechodzi go od końca, niosąc „sprzężenia” v̄ = ∂f/∂v. Jedno przejście wstecz daje pochodne jednego wyjścia względem wszystkich wejść. W uczeniu maszynowym wyjście jest jedno (strata), a wejść, czyli wag, są miliony, więc tryb odwrotny jest idealny. Propagacja wsteczna w sieciach to dokładnie tryb odwrotny AD.
Cena trybu odwrotnego to pamięć: trzeba przechowywać wartości pośrednie z przejścia w przód. Teoretyczny koszt czasowy jest ograniczony: gradient kosztuje stałą wielokrotność czasu obliczenia funkcji, typowo kilka razy więcej, niezależnie od liczby wejść.
PyTorch buduje graf dynamicznie, w trakcie wykonywania kodu (define-by-run). Każdy tensor z requires_grad=True zapamiętuje operację, która go stworzyła. Wywołanie backward() przechodzi ten graf od straty wstecz. Dzięki temu pętle i warunki w Pythonie działają naturalnie: graf jest tym, co faktycznie się wykonało w danym kroku. JAX wybiera inną drogę: przekształca funkcje (grad, jit, vmap) i łatwo łączy oba tryby.
Zastrzeżenia: AD różniczkuje program, a nie „matematyczną” funkcję. W punktach nieróżniczkowalnych (ReLU w zerze, abs, max) biblioteka przyjmuje umowną wartość pochodnej. Operacje dyskretne, takie jak argmax, zaokrąglanie czy losowanie, mają gradient zerowy lub nieokreślony, więc sygnał przez nie nie przepłynie.
Na przykładzie
Dla f(x, y) = x·y + sin(x) w punkcie x = 2, y = 3: v₁ = 6, v₂ = sin(2) ≈ 0,909, f ≈ 6,909. Tryb w przód z ziarnem (x' = 1, y' = 0) daje v₁' = 1·3 + 2·0 = 3 oraz v₂' = cos(2) ≈ −0,416, czyli ∂f/∂x ≈ 2,584. Aby dostać ∂f/∂y = 2, potrzebne jest drugie przejście z ziarnem (0, 1). Tryb odwrotny zaczyna od f̄ = 1, ustala v̄₁ = v̄₂ = 1, a potem jednym przejściem daje oba wyniki: x̄ = v̄₁·y + v̄₂·cos(x) ≈ 2,584 oraz ȳ = v̄₁·x = 2. PyTorch zwraca te same wartości.
Różniczkowanie numeryczne tej samej pochodnej pokazuje, czego AD unika. Błąd ilorazu różnicowego wynosi 0,045 dla h = 0,1, maleje do 4,5·10⁻⁶ dla h = 10⁻⁵ i do 3,2·10⁻⁸ dla h = 10⁻⁸. Dalej rośnie, bo zaczyna dominować błąd zaokrągleń: 1,4·10⁻⁴ dla h = 10⁻¹² i 0,81 dla h = 10⁻¹⁵, gdy przybliżenie wynosi 1,78 zamiast 2,58. Do tego dla sieci 64-64-10 na zbiorze Digits, z 4810 parametrami, gradient numeryczny wymagałby 4811 przejść w przód. Tryb odwrotny potrzebuje jednego przejścia w przód i jednego wstecz.
Dane: Digits (ręcznie pisane cyfry 8×8)
W praktyce
- PyTorch:
x = torch.tensor(2.0, requires_grad=True), obliczenie,y.backward(), wynik wx.grad. Funkcyjnie:torch.autograd.grad(y, [x]); tryb w przód:torch.func.jvp. torch.no_grad()itensor.detach()odcinają graf: używaj ich w ewaluacji i przy wartościach, które mają być traktowane jak stałe.- Gradienty w
.gradsumują się między wywołaniamibackward(); zeruj je (optimizer.zero_grad()), chyba że świadomie akumulujesz gradient z kilku batchy. - Operacje w miejscu (
x += 1,relu_) na tensorach potrzebnych do gradientu kończą się błędem albo cichym uszkodzeniem grafu; w razie wątpliwości unikaj ich. - Własną funkcję z ręcznym gradientem definiuje się przez
torch.autograd.Functioni weryfikujetorch.autograd.gradcheck.
Najczęstsze pytania
- Czy automatyczne różniczkowanie to to samo co propagacja wsteczna?
- Propagacja wsteczna jest szczególnym przypadkiem trybu odwrotnego AD, zastosowanym do sieci neuronowej i funkcji straty. AD jest ogólniejsze: różniczkuje dowolne programy, oferuje tryb w przód i pozwala liczyć pochodne wyższych rzędów.
- Dlaczego nie używać po prostu różnic skończonych?
- Są niedokładne, bo h musi być jednocześnie małe (błąd przybliżenia) i duże (błąd zaokrągleń), i wymagają osobnego obliczenia funkcji dla każdego parametru. Nadają się za to świetnie do testowania poprawności gradientów z AD.
- Czy AD potrafi liczyć drugie pochodne?
- Tak, przez zastosowanie AD do programu, który sam liczy gradient. W PyTorch wymaga to `create_graph=True` w `backward()` lub użycia `torch.func.hessian`. Pełny hesjan dla milionów parametrów jest jednak zbyt duży, więc zwykle liczy się iloczyny hesjanu z wektorem.
Źródła
- Baydin A. G., Pearlmutter B. A., Radul A. A., Siskind J. M., „Automatic differentiation in machine learning: a survey”, Journal of Machine Learning Research 18(153), 2018, s. 1–43.
- Griewank A., Walther A., „Evaluating Derivatives: Principles and Techniques of Algorithmic Differentiation”, 2nd ed., SIAM, 2008.
- Paszke A. i in., „PyTorch: An Imperative Style, High-Performance Deep Learning Library”, NeurIPS 2019.
- Goodfellow I., Bengio Y., Courville A., „Deep Learning”, MIT Press, 2016, podrozdz. 6.5.
- Dokumentacja PyTorch, „Automatic differentiation package — torch.autograd”: https://pytorch.org/docs/stable/autograd.html