ML Atlas

06 · Sieci · 4 min czytania · aktualizacja

Czym jest katastrofalne zapominanie i jak mu zapobiega replay?

W skrócie

Katastrofalne zapominanie to utrata starej wiedzy, gdy sieć trenuje się tylko na nowym zadaniu. Replay miesza stare przykłady z nowymi, by gradient dbał o oba.

Co to jest

Katastrofalne zapominanie (catastrophic forgetting, catastrophic interference) to gwałtowny spadek jakości na wcześniej nauczonym zadaniu, gdy sieć neuronową trenuje się dalej wyłącznie na nowym zadaniu lub na danych z nowego rozkładu. Zjawisko opisali McCloskey i Cohen w 1989 roku.

Najprostsza ochrona to replay (rehearsal): do nowych przykładów domieszkowuje się próbkę starych. Problem dotyczy wszystkich modeli uczonych gradientem na wspólnych wagach, w tym dostrajania dużych modeli językowych. Model trenowany od zera na nowych danych, np. drzewo, nie „zapomina” — po prostu nigdy nie miał starej wiedzy.

Mechanizm — dlaczego tak działa

Te same wagi niosą całą wiedzę sieci. Gradient liczony wyłącznie na nowych danych wskazuje kierunek, w którym spada strata na nowych danych, i nie zawiera żadnej informacji o tym, co dzieje się ze stratą na starych. Jeśli dobre rozwiązania obu zadań leżą w różnych obszarach przestrzeni wag, każdy krok przesuwa sieć w stronę nowego obszaru, „nie patrząc” na stary.

Po kilkudziesięciu krokach wynik na starym zadaniu spada, często do poziomu zgadywania lub niżej. To nie jest powolne zanikanie śladów pamięci, tylko skutek uboczny rozproszonej reprezentacji: w sieci nie ma osobnej szuflady na każde zadanie. Gdy nowe zadanie dotyczy nowych klas, dochodzi jeszcze warstwa wyjściowa: trening na samych nowych klasach uczy ją, że starych klas po prostu nie ma.

Replay zmienia funkcję straty: zamiast straty na nowych danych minimalizujemy stratę na mieszance starych i nowych. Gradient staje się sumą gradientów obu zadań, więc krok, który psułby stare zadanie, jest hamowany przez jego składnik. Cena jest realna: stare przykłady trzeba przechowywać, a kroki treningu dzielą się między zadania.

Alternatywy niewymagające starych danych to regularyzacja chroniąca wagi ważne dla starego zadania (EWC, Kirkpatrick i in. 2017 — kara proporcjonalna do informacji Fishera), destylacja z poprzedniej wersji modelu oraz osobne moduły lub adaptery na zadanie (np. LoRA), które zostawiają bazowe wagi nietknięte.

Zastrzeżenie: zapominanie jest tym silniejsze, im bardziej różnią się rozkłady zadań, im większy learning rate i im dłużej trwa trening tylko na nowym. Mały learning rate przy dostrajaniu to najtańsza, choć częściowa ochrona.

Na przykładzie

Zbiór Digits (cyfry 8×8 ze scikit-learn) podzieliliśmy na dwa zadania: A — cyfry 0–4, B — cyfry 5–9; 70% danych to trening (630 przykładów A, 627 B), 30% test. Sieć MLPClassifier z 64 neuronami ukrytymi (random_state=0) trenowana przez 30 epok na A rozpoznawała cyfry 0–4 na teście w 99%.

Potem trenowaliśmy ją dalej tylko na B. Po jednej epoce trafność na A spadła do 49%, po dwóch do 4%, po pięciu do 0% — sieć przestała wskazywać cyfry 0–4 w ogóle, choć B opanowała w 99%. Gdy do danych B domieszaliśmy losowe 63 przykłady z A (10% starego zbioru), po tych samych 30 epokach sieć zachowała 85% na A przy 99% na B; przy 157 przykładach (25%) — 88% na A.

Dane: Digits (ręcznie pisane cyfry 8×8)

W praktyce

  • Dostrajanie LLM na wąskim zbiorze instrukcji psuje umiejętności ogólne; standardem jest domieszka danych ogólnych (replay) i mały learning rate rzędu 10⁻⁵.
  • Uczenie ciągłe (continual learning): bufor pamięci z losową próbką starych przykładów; w RL to experience replay (Mnih i in. 2015).
  • PyTorch nie ma gotowego mechanizmu — mieszasz zbiory przez ConcatDataset lub własny sampler; biblioteka Avalanche implementuje replay, EWC i inne metody.
  • Mierz wynik na starym zadaniu po każdej fazie treningu; sam wynik na nowym nic nie mówi o zapominaniu.
  • Typowy błąd: douczanie modelu produkcyjnego wyłącznie na danych z ostatniego miesiąca — model przestaje rozpoznawać przypadki sezonowe.

Najczęstsze pytania

Co to jest katastrofalne zapominanie w sieciach neuronowych?
Utrata wcześniej nauczonych umiejętności podczas treningu na nowym zadaniu. Gradient nowego zadania przesuwa wspólne wagi bez względu na stare zadanie, więc jego wynik spada — często gwałtownie. Zjawisko opisali McCloskey i Cohen w 1989 roku.
Jak zapobiec zapominaniu przy dostrajaniu modelu?
Najprościej: mieszaj stare przykłady z nowymi (replay) i używaj małego learning rate. Bez dostępu do starych danych pomagają regularyzacja chroniąca ważne wagi (EWC), destylacja z poprzedniego modelu albo adaptery (LoRA), które zostawiają bazowe wagi bez zmian.
Czy replay ma wady?
Tak: trzeba przechowywać stare dane, co bywa problemem pamięci i prywatności, a kroki treningu dzielą się między zadania. W przykładzie na Digits sam replay nie przywrócił pełnych 99% na starym zadaniu, ale zamienił spadek do zera w spadek o kilkanaście punktów.

Źródła

  • McCloskey, M., Cohen, N. (1989). "Catastrophic interference in connectionist networks: the sequential learning problem". Psychology of Learning and Motivation 24, 109–165.
  • Robins, A. (1995). "Catastrophic forgetting, rehearsal and pseudorehearsal". Connection Science 7(2), 123–146.
  • Kirkpatrick, J. i in. (2017). "Overcoming catastrophic forgetting in neural networks". PNAS 114(13), 3521–3526. arXiv:1612.00796
  • Parisi, G., Kemker, R., Part, J., Kanan, C., Wermter, S. (2019). "Continual lifelong learning with neural networks: a review". Neural Networks 113, 54–71. arXiv:1802.07569
  • Mnih, V. i in. (2015). "Human-level control through deep reinforcement learning". Nature 518, 529–533 (experience replay).

Zobacz też