Перейти до вмісту
6/30Розділ 6 з 30

Як змусити мережу навчатися — і узагальнювати

Шар за шаром виправляємо мережу, чия втрата застигла на ln 2. Далі — double descent: 5 000 параметрів на 40 точках.

На цій сторінці

Мережа з розділу 5 працює. У неї дев’ять параметрів, вона навчається XOR, а її градієнти збігаються з PyTorch до шістнадцяти десяткових знаків.

Зробіть її шестишаровою — і вона повністю перестає навчатися. Не повільно — повністю. Ось шестишарова мережа для задачі класифікації двох спіралей, навчена протягом 5000 кроків:

TEXT
step    1: loss 0.693147
step 5000: loss 0.693147
accuracy: 50.0 %

Це число не випадкове. ln2=0.693147\ln 2 = 0.693147 — це бінарна крос-ентропія моделі, яка для всього видає ймовірність 0.50.5, а 50 % — це підкидання монети на збалансованому датасеті. Після п’яти тисяч кроків мережа не зрушила навіть на одну цифру. Нічого не впало, жодних попереджень, і градієнти все ще абсолютно правильні.

Цей розділ — про прірву між мережею, яка запускається, і мережею, яка працює. Він має дві половини, що виглядають як різні теми, але є однією роботою: змусити втрату йти вниз і змусити її йти вниз на даних, яких модель ніколи не бачила.

Почніть із погляду на факти, а не з припущень. Пропустіть batch входів крізь мережу й виведіть стандартне відхилення активацій на кожному шарі, а потім стандартне відхилення градієнтів ваг:

profile.pyPYTHON
def profile(model, x):
    h = x
    for layer in model:
        h = layer(h)
        if isinstance(layer, (nn.Tanh, nn.ReLU)):
            print(f"activation std: {h.std().item():.4f}")
    model(x).sum().backward()
    for p in model.parameters():
        if p.dim() == 2:
            print(f"gradient std: {p.grad.std().item():.2e}")

Три ініціалізації, та сама архітектура, шість шарів tanh\tanh:

ініціалізаціяstd активацій, шари 1→6
normal, std 0.010.010.0145 · 0.0016 · 0.0002 · 0.0000 · 0.0000 · 0.0000
normal, std 110.6573 · 0.9296 · 0.9585 · 0.9634 · 0.9637 · 0.9625
Xavier0.1579 · 0.1493 · 0.1353 · 0.1333 · 0.1325 · 0.1403
ініціалізаціяstd градієнтів, перший шар → останній
normal, std 0.010.013.20e-06 · 4.97e-07 · … · 6.40e-06
normal, std 111.94e+03 · 2.28e+02 · 1.22e+02 · 4.43e+01 · 1.85e+01 · 7.30e+00
Xavier2.31e+00 · 4.50e-01 · 4.26e-01 · 3.89e-01 · 4.39e-01 · 4.73e-01

Перший рядок — це мережа вище, і вона не навчається повільно: у ній не лишилося сигналу. До четвертого шару стандартне відхилення активацій у чотирьох десяткових знаках провалилося до нуля. Кожен вхід дає той самий вихід, вихід є константою, а градієнт константи — ніщо. Ваги ініціалізували малими «для безпеки», і ця малість виявилася фатальною.

Другий рядок — протилежний збій, і його варто зрозуміти, бо він контрінтуїтивний. Активації виглядають здоровими — близько 0.96 — але це tanh\tanh у насиченні, притиснута до своєї межі, саме той режим, який у розділі 5 був виміряний як втрата майже десяти тисяч разів у градієнті. І все ж градієнти величезні: 1940 на першому шарі. Обидві речі істинні одночасно. Кожен зворотний крок множить на WW^\top, а з 128 входами одиничної дисперсії цей множник має підсилення приблизно 12811\sqrt{128} \approx 11, що перекриває стискання від насиченого tanh\tanh. Градієнти геометрично зростають на шляху назад. Це вибуховий градієнт, і в реальному навчальному запуску він за кілька кроків дає значення втрати nan.

Третій рядок — те, що потрібно: активації мають приблизно сталий масштаб уздовж глибини, градієнти мають приблизно сталий масштаб уздовж глибини. Нічого не вмирає, нічого не вибухає.

Добра ініціалізація виправляє масштаб на нульовому кроці. Вона не тримає його фіксованим: ваги змінюються, і до п’ятитисячного кроку акуратний аргумент про дисперсію вже не діє.

Шари нормалізації підтримують масштаб постійно. Для вектора активацій: відніміть середнє, поділіть на стандартне відхилення, а потім застосуйте навчуваний масштаб γ\gamma і зсув β\beta, щоб шар міг скасувати нормалізацію, якщо виявиться, що саме цього він хоче:

h^=hμσ2+ϵ,y=γh^+β\hat{h} = \frac{h - \mu}{\sqrt{\sigma^2 + \epsilon}}, \qquad y = \gamma\hat{h} + \beta

Єдине справжнє питання — по чому усереднювати. Batch normalization3 бере μ\mu і σ\sigma уздовж batch-виміру, по одній статистиці на feature. Layer normalization4 бере їх уздовж features, по одній статистиці на приклад.

Цей вибір виглядає дрібним, але визначає майже все далі:

BatchNorm робить вихід кожного прикладу залежним від інших прикладів, які випадково опинилися в його batch. Під час навчання це м’який регуляризатор. Під час inference batch немає, тому доводиться зберігати ковзне середнє статистик, зібраних під час навчання, — а це означає, що шар поводиться по-різному в режимі навчання та оцінювання, і забути перемкнути режими — один із найпоширеніших багів у цій сфері. Він також гірше працює з малими batch і незручний для послідовностей змінної довжини, бо «середнє по batch у позиції 40» обчислюється з тієї кількості послідовностей, які випадково виявилися такими довгими.

LayerNorm нормалізує кожен приклад окремо. Немає залежності від batch, немає ковзних статистик, однакова поведінка під час навчання й inference, байдужість до розміру batch, байдужість до довжини послідовності. Кожна з цих властивостей стає не приємним бонусом, а вимогою, щойно ви генеруєте по одному token за раз для одного користувача — саме туди приводить розділ 13.

Саме тому LayerNorm знову зустрінеться вам у розділі 9 без змін: transformer block використовує його, і використовує з причин у правій колонці, а не тому, що він абстрактно «кращий».

Чотири кандидати на виправлення мертвої мережі: ініціалізація Xavier, LayerNorm, residual connections і Adam замість SGD. Спокуса — застосувати всі чотири й рухатися далі. Зробіть так — і ви ніколи не дізнаєтеся, що саме мало значення; наступного разу у вас не буде методу, лише ритуал.

Тому застосовуйте їх по одному. Той самий seed, ті самі дані, та сама архітектура, 800 кроків:

що доданофінальна втрататочність
нічого0.693150.0 %
ініціалізація Xavier0.569260.4 %
LayerNorm0.623061.5 %
residual connections0.665156.6 %
Adam0.678758.7 %
усі чотири0.0000100.0 %

Читайте цю таблицю так, як читали б її о 2-й ночі, і висновок буде: ніщо не працює окремо, усе працює разом, отже deep learning — алхімія. Цей висновок хибний, і з’ясувати чому — найкорисніше в цьому розділі.

Дайте кожному запуску вшестеро більший бюджет — 5000 кроків замість 800 — і картина повністю зміниться:

що доданофінальна втрата @ 5000точність
нічого0.693150.0 %
ініціалізація Xavier0.0007100.0 %
LayerNorm0.0002100.0 %
residual connections0.665356.7 %
Adam0.690853.4 %
Xavier + Adam0.0000100.0 %
Xavier + LayerNorm0.0001100.0 %

Тепер картина чітка, і це вже діагностика, а не ритуал.

Сама ініціалізація це виправляє. Сама нормалізація це виправляє. Кожна з них лікує справжню хворобу — колапс прямого сигналу до нуля — і будь-якої достатньо. На 800 кроках вони просто виглядали як частковий успіх, бо вже розв’язали проблему, але ще вибиралися назовні.

Residual connections і Adam не виправляють це за жодного бюджету. Не тому, що вони погані, а тому, що лікують іншу хворобу. Residual connection дає градієнту шлях в обхід заблокованого шару; це дуже цінно, коли проблема в градієнті, і нічого не варте, коли прямий сигнал уже нульовий, бо короткий шлях навколо мертвого шару все одно несе мертве значення. Adam масштабує крок кожного параметра за його власною історією градієнтів; це допомагає, коли градієнти мають дуже різні масштаби, але не може воскресити мережу, чий вихід не залежить від входу.

А «нічого» після п’яти тисяч кроків усе ще рівно 0.6931. Не 0.6929. Це не повільно; це мертво, і тепер цю різницю видно так, як раніше не було видно, бо є рядок із працюючим виправленням для порівняння.

Відтепер цей курс використовує PyTorch. Це має бути заслужено, а не просто оголошено, тож ось рівно те, що він робить із того, що ви вже вмієте робити.

Оптимізатор — це правило перетворення градієнтів на оновлення параметрів. Звичайний gradient descent використовує градієнт. Momentum використовує його ковзне середнє, що згладжує шум і набирає швидкість у напрямках, які лишаються узгодженими:

optim_by_hand.pyPYTHON
v = beta * v + p.grad          
p -= lr * v                    

Adam5 зберігає два ковзні середні — градієнта й градієнта у квадраті — і ділить одне на квадратний корінь з іншого, тож кожен параметр отримує крок, масштабований до його нещодавнього масштабу градієнта:

optim_by_hand.pyPYTHON
m = b1 * m + (1 - b1) * g          # mean of the gradient          
v = b2 * v + (1 - b2) * g * g      # mean of the squared gradient  
m_hat = m / (1 - b1 ** t)          # bias correction: both averages start at zero
v_hat = v / (1 - b2 ** t)
p -= lr * m_hat / (v_hat.sqrt() + eps)   

Десять рядків. Запустіть обидва проти torch.optim на тій самій задачі протягом 50 кроків:

TEXT
SGD+momentum   by hand [2.7781870365142822, -1.0304985046386719]
               torch   [2.7781870365142822, -1.0304983854293823]   max |diff| = 1.19e-07
Adam           by hand [0.4893140196800232, -0.46317872405052185]
               torch   [0.48931416869163513, -0.46317875385284424]   max |diff| = 1.49e-07

Ідентично до точності float32. torch.optim.Adam — це ті п’ять рядків плюс десятиліття турботи про крайові випадки й C++ kernel. Ось обмін, який ви робите відтепер: не магія замість розуміння, а швидкість за рядки, які ви вже написали.

Звичне пояснення Adam — «адаптивні learning rates для кожного параметра», але це опис, а не причина. Причина — геометрія, і її можна виміряти.

Візьміть втрату, чия кривина відрізняється між напрямками: крута в одному, полога в іншому. SGD має один глобальний learning rate, тож мусить вибрати значення, достатньо мале для стабільності в найкрутішому напрямку, — і це значення стає набагато замалим для пологого напрямку, де прогрес ледь повзе. Саме це спричиняє класичну картинку, де gradient descent зигзагом спускається вузькою долиною.

Два співвідношення кривини, три оптимізатори, 300 кроків, і кожному оптимізатору дали найкращий learning rate із перебору, щоб ніхто не мав фори чи штрафу:

співвідношення кривиниSGDSGD + momentumAdam
10 : 1помилка 0.000002помилка 0.000000помилка 0.000000
1000 : 1помилка 1.925485помилка 0.001432помилка 0.000000
розбігся на (1000:1)4 з 8 rates4 з 8 rates0 з 6 rates

За співвідношення десять працює все, і обговорювати нічого. За тисячі звичайний SGD не може дійти до відповіді за жодного випробуваного learning rate — його найкращий результат усе ще має помилку 1.93 — і прямо розбігається на половині rates. Adam точно потрапляє в target і не розбігається на жодному.

Остання колонка — практична причина, чому Adam є стандартом. Не тому, що Adam знаходить кращі рішення; на добре обумовлених задачах налаштований SGD часто зрівнюється з ним або перемагає. А тому, що Adam набагато менш чутливий до вибраного вами learning rate, а реальні мережі мають співвідношення кривини набагато гірші за тисячу серед мільйонів своїх параметрів.

Тут доречні ще дві речі, і обидві — в один рядок. Gradient clipping перемасштабує вектор градієнта щоразу, коли його норма перевищує поріг, перетворюючи рядок діагностичної таблиці «втрата раптом стрибає до величезного значення» на не-подію. І learning rate schedules: короткий warmup від майже нуля протягом перших кількох сотень кроків, бо оцінки дисперсії в Adam — сміття, доки вони не побачили кілька градієнтів, а повнорозмірний крок на смітті може зруйнувати ініціалізацію; потім cosine decay до нуля, бо завершувати запуск із тим самим розміром кроку, з яким ви почали, означає тремтіти навколо мінімуму, а не осідати в ньому.

Друга половина: модель, яка ідеально підганяється й нічого не передбачає

Посилання на розділ: Друга половина: модель, яка ідеально підганяється й нічого не передбачає

Досі все було про те, як змусити втрату знижуватися. Тепер складніша половина, бо зниження втрати — не мета; це proxy для мети, і цей proxy ламається конкретним і знаменитим способом.

Дванадцять точок із гладкої функції з невеликим шумом. Підженемо поліноми зростаючого степеня:

степіньtrain RMSEtest RMSE
10.7644990.6985
30.2526050.3031
50.1644370.1568
90.0889600.2347
110.0000001.2094

Степінь 11 через 12 точок проходить точно через кожну — train-помилка нуль до шести десяткових знаків — і на даних, яких вона не бачила, у вісім разів гірша за степінь 5. Попросіть степінь 3 і степінь 11 передбачити в x=3.25x = 3.25, просто за межами навчального діапазону:

TEXT
degree  3: predicts   -1.053   (truth -0.012)
degree 11: predicts  +61.224   (truth -0.012)

Шістдесят один, коли відповідь приблизно нуль. Модель не вивчила функцію; вона вивчила дванадцять точок, а між ними робить усе, чого вимагає арифметика.

Це overfitting, а його протилежність — степінь 1, який узагалі не може представити криву й поганий всюди, — це underfitting. Класичний опис розкладає очікувану помилку моделі на три частини: bias, помилка від того, що модель надто жорстка, щоб представити істину; variance, помилка від того, що модель настільки гнучка, що ганяється за шумом у цій конкретній вибірці; і незвідний шум, який не виправляє ніщо. Прості моделі мають bias, гнучкі моделі мають високу variance, а класичний рецепт — знайти золоту середину посередині: степінь 5 у таблиці вище.

Стандартні інструменти всі атакують складову variance:

  • L2-регуляризація (weight decay) додає λw2\lambda \lVert w \rVert^2 до втрати, притягуючи ваги до нуля й роблячи функцію гладкішою. У таблиці вище шкоди завдає найбільший коефіцієнт степеня 11; штраф за розмір його знешкоджує.
  • L1 натомість додає λwi\lambda \sum |w_i|. Різниця не косметична: градієнт L2 пропорційний вазі, тож зменшується разом із нею, наближаючись до нуля, але не доходячи до нього, тоді як градієнт L1 — це константа ±λ\pm\lambda, яка продовжує штовхати до кінця. Тому L1 дає ваги, що є точно нульовими, — вона відбирає features. L2 дає малі ваги. Використовуйте L2, коли потрібна гладкість, і L1, коли потрібна розрідженість.
  • Dropout7 зануляє випадкову підмножину активацій на кожному навчальному кроці, тож жоден юніт не може покладатися на присутність будь-якого конкретного іншого юніта.
  • Early stopping стежить за validation-втратою й зупиняє навчання, коли вона повертає вгору.
  • Data augmentation створює більше навчальних прикладів із тих, що вже є, атакуючи проблему в її джерелі: overfitting — це нестача даних не менше, ніж надлишок параметрів.
  • Cross-validation ділить дані kk способами й навчає kk разів, що купує надійну оцінку test-помилки, коли даних замало, щоб виділити окремий held-out набір.

Double descent, або чому попередній розділ — не вся історія

Посилання на розділ: Double descent, або чому попередній розділ — не вся історія

Тепер факт, який ламає картинку.

Історія bias-variance каже: після золотої середини більше параметрів означає гірше узагальнення. Сучасні language models мають набагато більше параметрів, ніж дозволяють класичні правила для даних, які вони бачать, і узагальнюють чудово. Обидва твердження істинні, і примирити їх — найкорисніша річ у цьому розділі.

Сорок навчальних точок, двадцятивимірні входи, випадкові ReLU features, і кількість features PP перебрано від 2 до 5000 — з вибором minimum-norm розв’язку щоразу, коли існує багато розв’язків, які підходять:

PPP/nP/ntrain RMSEtest RMSEw\lVert w \rVert
100.250.88221.25201.89
200.500.59621.16342.59
300.750.38961.53234.15
380.950.17693.716310.25
401.000.00005.814014.83
421.050.00003.16239.35
601.500.00001.10582.78
2005.000.00000.66380.98
150037.500.00000.58590.33
5000125.000.00000.56640.18

Читайте це у трьох частинах. До P/n=0.5P/n = 0.5 класична історія працює точно: помилка падає, а потім починає зростати. На P=n=40P = n = 40порозі інтерполяції, де модель має рівно стільки параметрів, щоб пройти через кожну навчальну точку, — test-помилка досягає піку, 5.81, уп’ятеро гірше за малу модель. Цей пік — класичне попередження, і він реальний.

Потім вона знову спускається. І продовжує спускатися: повз P=5nP = 5n, повз P=37nP = 37n, аж до P=125nP = 125n, де test-помилка 0.5664 краща за найкращий результат, якого будь-коли досягала недопараметризована модель. Модель із 5000 параметрами, підлаштована до 40 точок, — найкраща модель у таблиці.

Це double descent,89 і механізм видно в останній колонці. Щойно P>nP > n, існує нескінченно багато налаштувань параметрів, які точно підганяють навчальні дані, а те, яке саме ви отримаєте, залежить від вибору. Minimum-norm розв’язок обирає найменший, і w\lVert w \rVert показує, що це означає: він досягає піку 14.83 рівно на порозі — де є рівно один інтерполяційний розв’язок, і ви застрягли з ним, хоч би яким екстремальним він був, — а потім монотонно падає зі зростанням PP, бо більше параметрів означає більше інтерполяційних розв’язків на вибір, а отже найменший доступний стає меншим. На P=5000P = 5000 норма дорівнює 0.18, у вісімдесят разів менше, ніж на порозі.

Отже, додаткові параметри не додають складності. Вони додають вибір, а правило відбору витрачає цей вибір на простоту. Регуляризація не у функції втрат; вона в алгоритмі. Gradient descent з малої ініціалізації має задокументований bias до розв’язків із малою нормою, саме тому така поведінка проявляється в реальних мережах, навчених звичайним способом, а не лише в лінійній алгебрі вище.

Практичний наслідок, від якого залежить розділ 10: «у моделі більше параметрів, ніж даних, отже вона overfit» — невалідний аргумент. Це було хорошим правилом, коли моделі жили ліворуч від порога. Усе цікаве тепер живе далеко праворуч від нього, де правило розвертається.

Інструментів із цього розділу достатньо, щоб навчити мережу, яка працює на даних, що їх можна покласти в таблицю: рядки чисел, колонка міток.

Мова — не така. Перш ніж модель зможе передбачити наступне слово, щось має вирішити, що таке «слово» взагалі, — і відповідь не літери й не слова, а словник, який модель вивчає з сирих байтів навчальних даних. Це рішення, ухвалене один раз до початку навчання, визначає, скільки речей модель може сказати, скільки коштує запит і чому моделі, здатні скласти іспит із права, не можуть надійно порахувати літери у strawberry.

Розділ 7 будує tokenizer.


Для residual connections, використаних вище, див. He та ін., Deep Residual Learning for Image Recognition (arXiv:1512.03385). Building makemore Part 3: Activations & Gradients, BatchNorm Андрія Карпатого проходить діагностику гістограм активацій на реальній моделі й є найкращим практичним поясненням першої половини цього розділу. Лекції 8 і 11–13 з Learning From Data Ясера Абу-Мостафи належно викладають класичну теорію узагальнення, включно з частинами, які цей розділ стиснув до одного абзацу.

  1. Glorot, X. and Bengio, Y. Understanding the difficulty of training deep feedforward neural networks. AISTATS (2010). Аргумент про збереження дисперсії відтворено в рамці вище.

  2. He, K., Zhang, X., Ren, S. and Sun, J. Delving Deep into Rectifiers: Surpassing Human-Level Performance on ImageNet Classification. arXiv:1502.01852 (2015).

  3. Ioffe, S. and Szegedy, C. Batch Normalization: Accelerating Deep Network Training by Reducing Internal Covariate Shift. arXiv:1502.03167 (2015). Зауважте, що пояснення через «internal covariate shift» у назві згодом суттєво оскаржили; шар працює, але початкове пояснення чому — спірне.

  4. Ba, J. L., Kiros, J. R. and Hinton, G. E. Layer Normalization. arXiv:1607.06450 (2016).

  5. Kingma, D. P. and Ba, J. Adam: A Method for Stochastic Optimization. arXiv:1412.6980 (2014).

  6. Loshchilov, I. and Hutter, F. Decoupled Weight Decay Regularization. arXiv:1711.05101 (2017).

  7. Srivastava, N., Hinton, G., Krizhevsky, A., Sutskever, I. and Salakhutdinov, R. Dropout: A Simple Way to Prevent Neural Networks from Overfitting. JMLR 15, pp. 1929–1958 (2014).

  8. Belkin, M., Hsu, D., Ma, S. and Mandal, S. Reconciling modern machine-learning practice and the classical bias–variance trade-off. PNAS 116(32), pp. 15849–15854 (2019). Стаття, яка дала явищу назву.

  9. Nakkiran, P., Kaplun, G., Bansal, Y., Yang, T., Barak, B. and Sutskever, I. Deep Double Descent: Where Bigger Models and More Data Hurt. arXiv:1912.02292 (2019). Показує ефект у реальних глибоких мережах, причому як уздовж осі часу навчання, так і вздовж осі розміру моделі.

Готові довірити вибір моделі LIA?

Створюйте з усіма моделями ШІ в одному місці — почніть безкоштовно вже сьогодні.