Нейросеть на Go с нуля · Глава 7 из 10

Глава 7. Обучаем классификатор MNIST

Прогресс сохранится в этом браузере (войдите, чтобы синхронизировать).

Введение

Все части на месте: матрицы, слои, потери, обратное распространение и загрузчик MNIST. В этой главе мы соберём их в программу, которая обучает настоящую модель на 60 000 изображений и распознаёт рукописные цифры с точностью выше 97%.

Но обучить — половина дела. Вторая половина — понимать, что происходит: почему точность встала на 11%, почему loss стал NaN, почему на обучающей выборке 99%, а на валидационной 92%. Диагностика — главный навык этой главы, и он ценнее любой конкретной архитектуры.

Архитектура и её цена

Начнём с самой простой рабочей конфигурации: вход 784 (28×28 развёрнутые в вектор), один скрытый слой из 128 нейронов с ReLU, выход 10 с softmax.

net := &nn.Network{Layers: []nn.Layer{
    nn.NewDense(784, 128, nn.HeInit(rng)),
    nn.NewReLU(),
    nn.NewDense(128, 10, nn.XavierInit(rng)),
}}

Посчитаем параметры: первый слой — 784×128 весов плюс 128 смещений = 100 480; второй — 128×10 плюс 10 = 1290. Итого 101 770 обучаемых чисел. Каждый из них будет обновлён на каждом шаге, а шагов при батче 64 получается 938 на эпоху. За 20 эпох — почти два миллиарда обновлений параметров. Это ощутимая работа, и именно поэтому мы так возились с внутренним циклом умножения матриц во второй главе.

Softmax здесь не входит в сеть как слой: он объединён с кросс-энтропией в функции потерь (глава 4), а для предсказаний применяется отдельно. Так мы избегаем лишнего якобиана и численных проблем.

Почему именно 128

Число нейронов скрытого слоя — гиперпараметр, то есть то, что подбирается экспериментом, а не выводится. Ориентиры: 32 нейрона дают около 95%, 128 — около 97,5%, 512 — около 98%, но обучаются в четыре раза дольше и заметнее переобучаются. 128 — хорошая точка старта: результат уже приличный, обучение занимает минуты, а не десятки минут.

Цикл обучения с валидацией

Ключевое отличие от главы 5: после каждой эпохи мы считаем метрики на валидационной выборке — данных, которых сеть не видела.

type EpochStats struct {
    Epoch     int
    TrainLoss float64
    ValLoss   float64
    ValAcc    float64
    Elapsed   time.Duration
}

func trainEpoch(net *nn.Network, loss nn.Loss, ds *dataset.Dataset,
    cfg nn.TrainConfig, rng *rand.Rand) float64 {

    next := ds.Batches(cfg.BatchSize, rng)
    var total float64
    var batches int

    for {
        bx, by, ok := next()
        if !ok {
            break
        }
        logits := net.Forward(bx)
        total += loss.Value(logits, by)
        net.Backward(loss.Grad(logits, by))
        net.SGDStep(cfg.LR)
        batches++
    }
    return total / float64(batches)
}

// Evaluate считает потери и точность, НЕ трогая параметры.
func Evaluate(net *nn.Network, loss nn.Loss, ds *dataset.Dataset, batch int) (float64, float64) {
    var totalLoss float64
    var correct, seen, batches int

    for start := 0; start < ds.X.Rows; start += batch {
        end := min(start+batch, ds.X.Rows)
        bx := rowsRange(ds.X, start, end)
        by := rowsRange(ds.Y, start, end)

        logits := net.Forward(bx)
        totalLoss += loss.Value(logits, by)
        probs := nn.Softmax(logits)

        for i := 0; i < probs.Rows; i++ {
            if argmaxRow(probs, i) == argmaxRow(by, i) {
                correct++
            }
            seen++
        }
        batches++
    }
    return totalLoss / float64(batches), float64(correct) / float64(seen)
}

Валидацию гоняем батчами, а не всей выборкой сразу: матрица 6000×784, прошедшая через слой на 128 нейронов, создаёт вполне заметные промежуточные аллокации. И да — валидация не вызывает Backward и SGDStep: параметры при оценке не меняются никогда.

Что печатать в лог

fmt.Printf("эпоха %2d | train %.4f | val %.4f | acc %.2f%% | %v\n",
    s.Epoch, s.TrainLoss, s.ValLoss, s.ValAcc*100, s.Elapsed.Round(time.Millisecond))

Четыре числа — минимум, по которому диагностируется почти всё. Обе потери и точность пишем ещё и в CSV: сравнивать эксперименты по логу в терминале невозможно, а по таблице — легко.

Подбор learning rate

Learning rate — самый важный гиперпараметр. Если выбрать его неправильно, не спасут ни архитектура, ни данные.

Дешёвый и надёжный способ найти рабочий диапазон — прогнать одну-две эпохи с разными значениями и посмотреть на loss:

for _, lr := range []float64{0.001, 0.01, 0.05, 0.1, 0.5, 1.0} {
    net := buildNet(rand.New(rand.NewPCG(1, 2))) // один и тот же старт!
    l := trainEpoch(net, loss, train, nn.TrainConfig{BatchSize: 64, LR: lr}, rng)
    fmt.Printf("lr=%-6g loss после эпохи = %.4f\n", lr, l)
}

Типичная картина для нашей сети с чистым SGD: при 0,001 loss едва сдвинулся, при 0,1–0,5 упал сильнее всего, при 1,0 начал скакать. Берут значение чуть меньше того, при котором начинается нестабильность.

Важнейшая деталь эксперимента: сеть должна пересоздаваться с одним и тем же сидом. Иначе ты сравниваешь не learning rate, а разные случайные инициализации.

Ориентиры для чтения кривой

  • Loss почти не двигается — lr слишком мал (или градиент не доходит: проверь обратный проход).
  • Loss падает плавно и выполаживается — то, что нужно.
  • Loss прыгает вверх-вниз с большим размахом — lr великоват, попробуй в 3–5 раз меньше.
  • Loss растёт или стал NaN — lr слишком велик, уменьшай на порядок.

Ранняя остановка и лучшие веса

Обучение почти никогда не стоит гонять «до конца». В какой-то момент валидационная точность перестаёт расти и начинает падать — сеть переходит от обобщения к запоминанию. Ранняя остановка ловит этот момент.

type EarlyStopper struct {
    Patience int // сколько эпох терпим отсутствие улучшения

    bestAcc    float64
    bestParams [][]float64 // копия всех параметров лучшей эпохи
    since      int
}

// Update возвращает true, если пора останавливаться.
func (e *EarlyStopper) Update(net *nn.Network, valAcc float64) bool {
    if valAcc > e.bestAcc {
        e.bestAcc = valAcc
        e.bestParams = snapshot(net) // запоминаем ЛУЧШИЕ веса, а не последние
        e.since = 0
        return false
    }
    e.since++
    return e.since >= e.Patience
}

// snapshot делает глубокую копию всех параметров сети.
func snapshot(net *nn.Network) [][]float64 {
    var out [][]float64
    for _, l := range net.Layers {
        params, _ := l.Params()
        for _, p := range params {
            out = append(out, append([]float64(nil), p.Data...))
        }
    }
    return out
}

Обрати внимание на append([]float64(nil), p.Data...): без копирования ты сохранишь не снимок, а ссылку на живой срез, который продолжит меняться при обучении, и «лучшие веса» окажутся последними. Это ровно та же ловушка, что с буфером картинок в главе 6.

Значение Patience порядка 3–5 эпох — разумный старт. Ставить 1 не стоит: валидационная точность естественно колеблется, и одна неудачная эпоха ещё ничего не значит.

Матрица ошибок

Точность 97% — одно число, и оно скрывает структуру ошибок. Матрица ошибок (confusion matrix) показывает, что с чем путается: строка — истинный класс, столбец — предсказанный.

// ConfusionMatrix[i][j] — сколько раз класс i был предсказан как j.
func ConfusionMatrix(net *nn.Network, ds *dataset.Dataset, classes int) [][]int {
    cm := make([][]int, classes)
    for i := range cm {
        cm[i] = make([]int, classes)
    }
    probs := nn.Softmax(net.Forward(ds.X))
    for i := 0; i < probs.Rows; i++ {
        cm[argmaxRow(ds.Y, i)][argmaxRow(probs, i)]++
    }
    return cm
}

Напечатай её после обучения. На MNIST ты почти наверняка увидишь знакомую картину: 4 путается с 9, 3 с 5 и 8, 7 с 1. Это не дефект модели, а свойство данных — эти цифры действительно похожи в рукописном виде, и часть примеров не разобрал бы и человек.

Практическая ценность матрицы в том, что она подсказывает, куда смотреть. Если один класс распознаётся заметно хуже остальных, причина обычно в данных (мало примеров, нестандартное написание), а не в архитектуре. Если же сеть валит всё в один класс — это симптом полностью сломанного обучения.

Диагностика: четыре типичных сценария

Loss стал NaN

Порядок действий: уменьшить lr в десять раз; проверить, что в кросс-энтропии есть eps внутри логарифма; проверить, что softmax вычитает максимум. В 95% случаев виновата одна из этих трёх причин. Полезно добавить в цикл обучения проверку math.IsNaN(loss) с немедленной остановкой и сообщением — иначе программа честно докрутит 20 эпох, обучая NaN.

Точность стоит на 11%

11% — это доля самого частого класса в MNIST, то есть сеть выдаёт одну и ту же цифру для всего. Причины: нулевая инициализация весов, ошибка знака в шаге обновления (+= вместо -=), обратный проход в неправильном порядке или градиент, который вообще не доходит до первого слоя. Проверь GradCheck — он найдёт первые три за минуту.

Train 99%, val 92%

Переобучение: сеть запомнила обучающие примеры. Лечится регуляризацией (глава 8), уменьшением сети или увеличением данных. Разрыв в один-два процентных пункта — норма; в семь — повод действовать.

И train, и val около 90% и не растут

Недообучение: модели не хватает выразительности или обучение остановилось слишком рано. Увеличь скрытый слой, добавь второй, обучай дольше, проверь lr. В отличие от переобучения, здесь бесполезно бороться с регуляризацией — её надо, наоборот, ослабить.

Кейс: полная программа обучения

func main() {
    dataDir := flag.String("data", "data/mnist", "каталог с файлами MNIST")
    epochs := flag.Int("epochs", 20, "число эпох")
    lr := flag.Float64("lr", 0.1, "скорость обучения")
    batch := flag.Int("batch", 64, "размер батча")
    seed := flag.Uint64("seed", 42, "сид генератора")
    out := flag.String("out", "model.bin", "куда сохранить веса")
    flag.Parse()

    train, val, test, err := dataset.Load(*dataDir)
    if err != nil {
        log.Fatalf("данные: %v", err)
    }
    log.Printf("обучение: %d, валидация: %d, тест: %d",
        train.X.Rows, val.X.Rows, test.X.Rows)

    rng := rand.New(rand.NewPCG(*seed, *seed+1))
    net := &nn.Network{Layers: []nn.Layer{
        nn.NewDense(784, 128, nn.HeInit(rng)),
        nn.NewReLU(),
        nn.NewDense(128, 10, nn.XavierInit(rng)),
    }}
    loss := nn.NewSoftmaxCrossEntropy()
    stopper := &EarlyStopper{Patience: 4}

    for epoch := 1; epoch <= *epochs; epoch++ {
        start := time.Now()
        trainLoss := trainEpoch(net, loss, train, nn.TrainConfig{
            BatchSize: *batch, LR: *lr,
        }, rng)

        if math.IsNaN(trainLoss) {
            log.Fatal("loss стал NaN: уменьши lr или проверь softmax/кросс-энтропию")
        }

        valLoss, valAcc := Evaluate(net, loss, val, 256)
        log.Printf("эпоха %2d | train %.4f | val %.4f | acc %.2f%% | %v",
            epoch, trainLoss, valLoss, valAcc*100, time.Since(start).Round(time.Millisecond))

        if stopper.Update(net, valAcc) {
            log.Printf("ранняя остановка: %d эпох без улучшения", stopper.Patience)
            break
        }
    }

    stopper.Restore(net) // возвращаем ЛУЧШИЕ веса, а не последние
    _, testAcc := Evaluate(net, loss, test, 256)
    log.Printf("итоговая точность на тесте: %.2f%%", testAcc*100)

    if err := nn.Save(*out, net); err != nil { // формат разберём в главе 10
        log.Fatalf("сохранение: %v", err)
    }
}

Типичный вывод: после первой эпохи около 93%, после третьей — 96%, к десятой — 97,5%, дальше рост замедляется до долей процента. Одна эпоха на обычном ноутбуке занимает несколько секунд. Тестовую выборку мы трогаем ровно один раз, в самом конце, — и это единственное честное число во всём логе.

Все гиперпараметры вынесены во флаги неслучайно: экспериментировать перекомпиляцией с правкой констант — верный способ через час не помнить, какой результат каким кодом получен. С флагами команда запуска сама себя документирует и попадает в лог.

Типичные ошибки

1. Оценивать на обучающей выборке

Точность на данных, которые сеть видела, всегда выше и ничего не говорит об обобщении. Это не педантизм: разница между 99% на обучении и 92% на валидации — реальная разница между «модель работает» и «модель бесполезна».

2. Сохранять последние веса вместо лучших

Если после пика точности сеть ещё пять эпох ухудшалась, а ты сохранил финальные веса, вся ранняя остановка теряет смысл. Снимок лучшей эпохи и восстановление перед сохранением — обязательны.

3. Подбирать гиперпараметры по тесту

Прогнать десять конфигураций, выбрать лучшую по тесту и отчитаться его точностью — значит завысить оценку. Тест перестаёт быть независимым, как только ты принял по нему хоть одно решение.

4. Менять сид вместе с гиперпараметром

Сравнение теряет смысл: разница может быть эффектом другой инициализации. Сид фиксируется, меняется ровно один параметр.

5. Продолжать обучение при NaN

После первого NaN все веса становятся NaN, и следующие эпохи — потерянное время. Проверка в цикле обучения и немедленный выход экономят минуты на каждом неудачном запуске.

6. Судить по одной эпохе

Первая эпоха бывает обманчива: конфигурация, выигравшая на старте, к десятой эпохе может проиграть. Для быстрого перебора lr одной эпохи достаточно, для выбора архитектуры — нет.

Практика

Задание 1. Полный цикл обучения

Собери cmd/train/main.go с флагами -data, -epochs, -lr, -batch, -seed, -out. Обучи сеть 784-128-10 и добейся точности выше 97% на валидации. Все метрики пиши и в лог, и в CSV.

Задание 2. Поиск learning rate

Прогони одну эпоху при lr = 0.001, 0.01, 0.05, 0.1, 0.5, 1.0 с фиксированным сидом. Построй таблицу «lr → loss после эпохи» и выбери рабочее значение. Отдельно найди значение, при котором обучение разваливается, и запиши, как выглядел лог в этот момент.

Задание 3. Ранняя остановка

Реализуй EarlyStopper с методами Update и Restore. Проверь, что после восстановления точность на валидации совпадает с лучшей зафиксированной в логе — если нет, ты сохранил ссылку вместо копии.

Задание 4. Матрица ошибок

Посчитай и напечатай матрицу ошибок на тесте с выравниванием по столбцам. Найди три самые частые пары путаницы и выведи для одной из них десять ошибочных примеров ASCII-визуализацией из главы 6. Посмотри на них: часть ошибок ты и сам бы допустил.

Задание 5. Влияние размера скрытого слоя

Обучи сети с 16, 32, 128 и 512 нейронами при одинаковых lr, батче и сиде. Составь таблицу: точность на валидации, время эпохи, разрыв между train и val. Найди, с какого размера начинается заметное переобучение.

Задание 6. Вторая архитектура

Добавь второй скрытый слой (784-256-128-10) и сравни с однослойной сетью. Не удивляйся, если выигрыш окажется скромным: на MNIST полносвязные сети упираются примерно в 98%, и дальше нужны свёртки — о них глава 9.

Чек-лист самопроверки

  • Точность на валидации выше 97%, на тесте — не более чем на полпроцента ниже.
  • Loss на обучении монотонно убывает, кривая гладкая.
  • Валидационные метрики считаются без вызова Backward и обновления весов.
  • Ранняя остановка срабатывает и восстанавливает именно лучшие веса.
  • Матрица ошибок построена; частые путаницы (4/9, 3/5, 7/1) видны.
  • Повтор запуска с тем же сидом даёт ту же точность до сотых.
  • Тестовая выборка использована ровно один раз.

Итог

Полносвязная сеть 784-128-10 обучается на MNIST до 97–98% за считаные минуты на CPU. Цикл обучения обязан включать валидацию после каждой эпохи, раннюю остановку с сохранением лучших весов и запись метрик в файл. Learning rate подбирается коротким экспериментом при фиксированном сиде, а loss, ставший NaN, означает одно из трёх: большой lr, отсутствие eps в логарифме или softmax без вычитания максимума.

Мы обучали чистым SGD с постоянным шагом — самым простым из возможных оптимизаторов. В следующей главе посмотрим, что даёт инерция, адаптивный шаг и Adam, и как бороться с переобучением через L2 и dropout.

Содержание серии (10)