Machine Learning Academy · Lekcja

Problem zanikającego gradientu w głębokich krokach czasowych

Uczą się Państwo obserwować eksplodujące i zanikające gradienty w głębokiej sieci RNN dzięki rejestrowaniu norm gradientów oraz rozumieć, dlaczego długie sekwencje destabilizują trenowanie.

Lekcja 2 z 413 kroki

Problem zanikającego gradientu w głębokich krokach czasowych to bezpłatna lekcja Machine Learning Academy na CoddyKit. To lekcja 2 z 4. Możesz przeczytać całą lekcję poniżej za darmo — a potem ćwiczyć ją interaktywnie w przeglądarce z wbudowanym edytorem kodu i tutorem AI dostępnym 24/7. To część ścieżki edukacyjnej Machine Learning Academy, a Twój postęp synchronizuje się między webem a aplikacją CoddyKit. Kurs Machine Learning Academy zawiera 4 lekcji w sumie.

Gradienty muszą przechodzić przez czas

Aby uczyć się zależności dalekiego zasięgu w sekwencji, gradienty pochodzące od funkcji straty w ostatnim kroku czasowym muszą przejść wstecz przez każdy krok czasowy, aby zaktualizować parametry przetwarzające początkowe dane wejściowe. Dla sekwencji o długości T oznacza to mnożenie tej samej macierzy wag W_hh przez samą siebie T razy podczas BPTT. To wielokrotne mnożenie jest przyczyną zarówno zanikających gradientów (wykładniczy zanik), jak i eksplodujących gradientów (wykładniczy wzrost).

import torch

# Conceptual illustration of gradient travel through T steps
# Gradient = dL/dh_T * (W_hh)^T * ...

# If W_hh has spectral radius < 1:
W_small = torch.eye(4) * 0.9
print('W^10 max value:', (W_small @ W_small @ W_small @
      W_small @ W_small @ W_small @
      W_small @ W_small @ W_small @ W_small).abs().max().item())
# -> very small: gradient vanishes

# If W_hh has spectral radius > 1:
W_big = torch.eye(4) * 1.1
print('W^10 max value:', (W_big ** 10).abs().max().item())
# -> very large: gradient explodes

Zanikający gradient: matematyczna przyczyna

Podczas BPTT gradient funkcji straty względem stanu ukrytego w kroku czasowym t obejmuje iloczyn macierzy Jacobiego pochodnych h względem h w każdym kroku od t do T. Macierz Jacobiego w każdym kroku zawiera wyrażenie diag(f'(h_t)) * W_hh, gdzie f' jest pochodną funkcji aktywacji. Dla tanh wartość f' jest ograniczona przez 1, a typowe losowe wagi mają promień spektralny mniejszy niż 1 — dlatego iloczyn T macierzy sprowadza gradienty do zera w tempie wykładniczym wraz ze wzrostem T.

import torch

# Track gradient norm through BPTT
def simulate_bptt_gradient(T, weight_scale=0.9):
    W = torch.eye(8) * weight_scale
    grad = torch.ones(8)   # gradient at final timestep
    norms = [grad.norm().item()]
    for t in range(T):
        grad = W.T @ grad  # one BPTT step
        norms.append(grad.norm().item())
    return norms

norms = simulate_bptt_gradient(T=20)
print('Gradient norms over 20 steps:')
print([f'{n:.4f}' for n in norms[::5]])
# Decreases from 2.83 -> nearly 0 after 20 steps

Eksperymentalna obserwacja zanikających gradientów

Zanikające gradienty można bezpośrednio zaobserwować, rejestrując normę gradientu w każdym kroku czasowym podczas BPTT. Należy zarejestrować haki wsteczne dla stanu ukrytego każdego kroku RNN, aby przechwycić wartości gradientów. W zwykłej sieci RNN składającej się z 50 kroków gradient w kroku 1 będzie zwykle równy 1e-6 lub mniejszy — praktycznie zerowy — co oznacza, że kilka pierwszych tokenów sekwencji ma niemal żadnego wpływu na parametry modelu. Model nie może nauczyć się, że podmiot na początku długiego zdania określa orzeczenie na jego końcu.

import torch
import torch.nn as nn

rnn = nn.RNN(4, 8, batch_first=True)
X = torch.randn(1, 30, 4, requires_grad=True)

output, h_n = rnn(X)
loss = output[:, -1, :].sum()   # loss at last timestep
loss.backward()

# Gradient with respect to early inputs
if X.grad is not None:
    per_step_grads = X.grad.abs().mean(dim=-1)
    print('Gradient norms per timestep (first 5 vs last 5):')
    print(per_step_grads[0, :5].tolist())    # early: tiny
    print(per_step_grads[0, -5:].tolist())   # late: larger

Eksplodujące gradienty: druga skrajność

Eksplodujące gradienty występują, gdy promień spektralny W_hh przekracza 1 — normy gradientów rosną wykładniczo wraz z długością sekwencji. Objawami są wartości straty NaN lub parametry aktualizowane do nieskończoności. W przeciwieństwie do zanikających gradientów (które powodują niezauważalne niepowodzenie uczenia) eksplodujące gradienty w widoczny sposób przerywają uczenie. Standardowym rozwiązaniem jest obcinanie gradientu: przed wykonaniem kroku optymalizatora należy przeskalować wektor gradientu tak, aby jego maksymalna norma L2 wynosiła 1.0. Zapobiega to katastrofalnym aktualizacjom bez usuwania sygnału gradientu.

import torch
import torch.nn as nn
import torch.optim as optim

rnn = nn.RNN(4, 8, batch_first=True)
optimizer = optim.SGD(rnn.parameters(), lr=0.1)

X = torch.randn(2, 50, 4)   # 50-step sequence
output, _ = rnn(X)
loss = output.sum()
loss.backward()

# Check gradient norm before clipping
total_norm = 0
for p in rnn.parameters():
    if p.grad is not None:
        total_norm += p.grad.data.norm(2) ** 2
total_norm = total_norm ** 0.5
print(f'Gradient norm before clip: {total_norm:.2f}')

# Clip to max_norm=1.0
nn.utils.clip_grad_norm_(rnn.parameters(), max_norm=1.0)
optimizer.step()

Wizualizacja norm gradientów w kolejnych warstwach

Praktyczną techniką debugowania jest rejestrowanie norm gradientów wszystkich parametrów po każdym przebiegu wstecznym i wykreślanie ich w trakcie uczenia. W zwykłych sieciach RNN macierz wag rekurencyjnych W_hh zwykle wykazuje znacznie mniejsze gradienty niż wagi wejściowe W_xh, co potwierdza, że informacje dalekiego zasięgu nie docierają do wcześniejszych parametrów. Taka wizualizacja często ujawnia, że tylko kilka ostatnich kroków czasowych ma istotny udział w uczeniu, co uzasadnia przejście na architektury z bramkami.

import torch
import torch.nn as nn

rnn = nn.RNN(4, 8, batch_first=True, num_layers=1)
X = torch.randn(1, 20, 4)
out, _ = rnn(X)
out.sum().backward()

print('Gradient norms per parameter:')
for name, p in rnn.named_parameters():
    if p.grad is not None:
        norm = p.grad.norm().item()
        print(f'  {name}: {norm:.6f}')
# weight_ih_l0 (input weights): larger
# weight_hh_l0 (recurrent weights): often much smaller

Dlaczego tanh pogłębia problem zanikania

Funkcja aktywacji tanh przyjmuje wartości z przedziału od -1 do 1, a jej pochodną jest 1 - tanh^2(x). Gdy dane wejściowe mają dużą wartość (funkcja jest nasycona), pochodna zbliża się do 0, niemal całkowicie odcinając gradient w tym kroku. Mnożenie wielu prawie zerowych pochodnych podczas BPTT nasila problem zanikania. ReLU ma pochodną równą 1 dla dodatnich danych wejściowych (nie występuje nasycenie), co ułatwia przepływ gradientu w sieciach jednokierunkowych, ale w sieciach RNN nadal dominuje wielokrotne mnożenie przez W_hh, które przy ReLU może powodować eksplozję gradientów.

import torch

# Tanh derivative: 1 - tanh(x)^2
x = torch.linspace(-4, 4, 9)
tanh_x = torch.tanh(x)
tanh_deriv = 1 - tanh_x ** 2

print('x:         ', x.tolist())
print('tanh(x):   ', [f'{v:.2f}' for v in tanh_x.tolist()])
print('tanh_deriv:', [f'{v:.2f}' for v in tanh_deriv.tolist()])
# At x=+/-3: deriv ~0.01 -- 100x smaller than at x=0
# Multiplied over 20 steps: 0.01^20 = 10^-40!

Skrócone BPTT: praktyczne obejście problemu

Skrócone BPTT ogranicza propagację gradientu do stałego okna K kroków czasowych zamiast do całej długości sekwencji. Gradienty są propagowane wstecz o K kroków, a następnie stan ukryty zostaje odłączony od grafu obliczeń (stając się stałą). Zapobiega to problemom z pamięcią i eksplozją gradientów dla bardzo długich sekwencji (dźwięku, korpusów tekstowych), ale uniemożliwia uczenie się zależności obejmujących więcej niż K kroków. W modelowaniu języka za pomocą zwykłych sieci RNN typowa wartość K wynosi 20–50.

import torch
import torch.nn as nn

rnn = nn.RNN(4, 8, batch_first=True)
batch_size = 4
h = torch.zeros(1, batch_size, 8)  # initial hidden state

# Process a 200-step sequence in chunks of 20
full_sequence = torch.randn(batch_size, 200, 4)

for chunk_start in range(0, 200, 20):
    chunk = full_sequence[:, chunk_start:chunk_start+20, :]
    out, h = rnn(chunk, h.detach())  # detach: stop grad here
    loss = out.sum()
    loss.backward()
    print(f'Chunk {chunk_start}-{chunk_start+20}: done')

Wyzwanie związane z zależnościami dalekiego zasięgu

Rozważmy zdanie: „The trophy that the man won was big.” Czasownik „was” musi zgadzać się z „trophy”, a nie z „man”. Wymaga to przeniesienia informacji o „trophy” przez 5 słów do miejsca, w którym pojawia się „was”. Zwykła sieć RNN uczona za pomocą BPTT praktycznie nie potrafi robić tego niezawodnie dla przerw dłuższych niż 5–10 tokenów. Jest to podstawowe ograniczenie, które doprowadziło do opracowania LSTM (1997), a później Transformerów (2017); oba rozwiązania mają jawne mechanizmy podtrzymywania informacji dalekiego zasięgu.

# Classic long-range dependency examples:
examples = [
    'The trophy ... man ... was [big/big] -- which subject?',
    'The cat ... [sat/sat] -- past vs present?',
    'The key [was/were] -- singular subject far away'
]

for ex in examples:
    print('Example:', ex)

# Vanilla RNN performance on long-range deps:
print('\nVanishing gradient effect on learning:')
for gap in [1, 5, 10, 20, 50]:
    ability = 'easy' if gap < 5 else ('hard' if gap < 20 else 'nearly impossible')
    print(f'  {gap}-step gap: {ability} for vanilla RNN')

Sposoby inicjalizacji wag w sieciach RNN

kilka sposobów inicjalizacji poprawia uczenie zwykłych sieci RNN na sekwencjach o umiarkowanej długości. Zainicjalizowanie W_hh jako macierzy ortogonalnej (o promieniu spektralnym dokładnie równym 1) zapobiega początkowemu zanikaniu i eksplozji gradientów. Dodanie połączenia pomijającego z wejścia bezpośrednio do wyjścia omija kilka mnożeń macierzy. Wykazano, że inicjalizacja macierzą jednostkową W_hh z aktywacją ReLU (IRNN) dorównuje LSTM w niektórych zadaniach, dowodząc, że sama inicjalizacja może częściowo rozwiązać problem zanikającego gradientu.

import torch
import torch.nn as nn

rnn = nn.RNN(4, 8, batch_first=True)

# Orthogonal init for hidden-to-hidden weights
nn.init.orthogonal_(rnn.weight_hh_l0)

# Identity init (IRNN) for W_hh
nn.init.eye_(rnn.weight_hh_l0)  # identity matrix

print('Spectral radius after orthogonal init:')
eigvals = torch.linalg.eigvals(rnn.weight_hh_l0)
print(eigvals.abs().max().item())  # should be ~1.0

Dlaczego wynaleziono LSTM

Problem zanikającego gradientu w sieciach RNN został opisany przez Hochreitera w 1991 roku. Jego rozwiązanie, sieć Long Short-Term Memory (LSTM) wprowadzona w 1997 roku, zastępuje pojedynczy stan ukryty stanem komórki chronionym przez bramki. Stan komórki przepływa w czasie z użyciem jedynie modyfikacji addytywnych (a nie multiplikatywnych), tworząc autostradę gradientu, która pozwala gradientom przepływać wstecz bez końca, bez zanikania. Ta pojedyncza innowacja architektoniczna umożliwiła praktyczne uczenie sekwencji z zależnościami obejmującymi ponad 100 kroków czasowych.

# The core difference between RNN and LSTM gradient flow:

# Vanilla RNN: h_t = tanh(W_hh * h_{t-1} + W_xh * x_t)
# Gradient must pass through tanh and W_hh MULTIPLICATIVELY
# -> vanishes after ~10 steps

# LSTM: c_t = f_t * c_{t-1} + i_t * g_t
# Cell state c_t is updated ADDITIVELY
# Forget gate f_t can be near 1 (keep everything)
# -> gradient flows back cleanly

print('LSTM key insight: additive cell state update')
print('Gradient highway: constant error carousel')
print('Forget gate f_t controls information retention')

Porównanie stabilności uczenia RNN i LSTM

Różnica w stabilności uczenia między zwykłymi sieciami RNN a LSTM staje się wyraźna dla sekwencji dłuższych niż 20–30 kroków czasowych. W klasycznym zadaniu kopiowania (odtworzeniu sekwencji wejściowej po długim opóźnieniu) zwykłe sieci RNN całkowicie zawodzą przy opóźnieniach większych niż 10 kroków, podczas gdy LSTM radzi sobie z opóźnieniami przekraczającymi 100 kroków. Ten praktyczny test jednoznacznie pokazuje, że problem zanikającego gradientu fundamentalnie ogranicza zwykłe sieci RNN oraz że architektoniczne rozwiązanie zastosowane w LSTM jest niezbędne w rzeczywistym modelowaniu sekwencji.

import torch
import torch.nn as nn

# Compare RNN vs LSTM on a 30-step sequence
models = {
    'RNN':  nn.RNN(4, 16, batch_first=True),
    'LSTM': nn.LSTM(4, 16, batch_first=True)
}

X = torch.randn(8, 30, 4)  # 30-step sequence

for name, model in models.items():
    out, _ = model(X)
    loss = out.sum()
    loss.backward()
    # Check gradient of first input vs last input
    total_grad_norm = sum(
        p.grad.norm().item() for p in model.parameters()
        if p.grad is not None
    )
    print(f'{name} total grad norm: {total_grad_norm:.4f}')

Szybkie sprawdzenie

Sprawdź swoją wiedzę na temat koncepcji uczenia maszynowego w języku Python przedstawionych w tej lekcji.

Podsumowanie lekcji

W tej lekcji dowiedział się Pan / dowiedziała się Pani, że: zanikające gradienty występują, gdy wielokrotne mnożenie przez W_hh (o promieniu spektralnym < 1) prowadzi do wykładniczego zaniku gradientów w długich sekwencjach, eksplodujące gradienty występują, gdy promień spektralny > 1, a rozwiązaniem jest obcinanie gradientu, oraz że LSTM wynaleziono specjalnie po to, aby rozwiązać problem zanikającego gradientu za pomocą addytywnej aktualizacji stanu komórki, która zapewnia autostradę gradientu. W następnej części szczegółowo przyjrzymy się architekturze komórki LSTM.

Bezpłatny start

Ucz się Python dzięki korepetycjom AI — za darmo

Pisz i uruchamiaj kod w przeglądarce, otrzymuj natychmiastową pomoc od korepetytora AI dostępnego 24/7 i kontynuuj naukę w sieci lub w aplikacji.

Kursy
30
Lekcje
120

Często zadawane pytania

Czy lekcja „Problem zanikającego gradientu w głębokich krokach czasowych” jest bezpłatna?

Tak — pełny tekst „Problem zanikającego gradientu w głębokich krokach czasowych” jest dostępny za darmo tutaj w sieci. Aby ćwiczyć ją interaktywnie (wbudowany edytor kodu i tutor AI dostępny 24/7) i odblokować resztę kursu Machine Learning Academy, przejdź na CoddyKit PRO. Kurs Machine Learning Academy zawiera 4 lekcji w sumie.

Co nauczysz się w „Problem zanikającego gradientu w głębokich krokach czasowych”?

Uczą się Państwo obserwować eksplodujące i zanikające gradienty w głębokiej sieci RNN dzięki rejestrowaniu norm gradientów oraz rozumieć, dlaczego długie sekwencje destabilizują trenowanie. Ćwiczysz Machine Learning Academy z praktycznym kodem, który uruchamiasz bezpośrednio w przeglądarce, a tutor AI dostępny 24/7 odpowiada na Twoje pytania podczas pracy nad lekcją.

Czy potrzebuję doświadczenia, aby zacząć Machine Learning Academy?

Nie wymagamy żadnego doświadczenia. Machine Learning Academy w CoddyKit jest strukturyzowany dla początkujących i zaawansowanych użytkowników, więc możesz zacząć tutaj lub od początku i uczyć się w swoim tempie. To lekcja 2 z 4.

Ile czasu zajmuje lekcja „Problem zanikającego gradientu w głębokich krokach czasowych”?

Większość lekcji CoddyKit trwa około 5–10 minut. Każda lekcja to mały, interaktywny krok, dzięki czemu robisz systematyczne postępy i zawsze wracasz dokładnie do tego samego miejsca — na webie i w aplikacji.

Czy mogę pisać i uruchamiać kod w tej lekcji Machine Learning Academy?

Tak. Każda lekcja Machine Learning Academy zawiera wbudowany edytor kodu, więc piszesz i uruchamiasz prawdziwy kod bezpośrednio w przeglądarce i od razu otrzymujesz sprzężenie zwrotne od AI — bez konfiguracji na komputerze.

Wszystkie lekcje w tym kursie

  1. Klasyczne sieci RNN: stan ukryty i rozwijanie sekwencji
  2. Problem zanikającego gradientu w głębokich krokach czasowych
  3. Komórka LSTM: bramki wejścia, zapominania i wyjścia
  4. Sequence-to-One: analiza sentymentu za pomocą LSTM
← Powrót do Machine Learning Academy