Глава 6. MNIST: загрузка настоящих данных на Go

214 просмотров
0 лайков
0 в избранном
Войдите, чтобы поставить лайк. Лайков:

Введение

Спирали были удобны: генерируются кодом, всегда под рукой. Настоящие данные так себя не ведут — они лежат в файлах непонятного формата, весят десятки мегабайт, содержат сюрпризы и требуют подготовки. Работа с данными занимает в реальных проектах больше времени, чем сама модель, и глава об этом в курсе обязана быть.

MNIST — 70 000 картинок рукописных цифр 28×28 пикселей в оттенках серого, стандартный «Hello, world» машинного обучения. Он хранится в формате IDX — простом двоичном формате, который мы разберём байт за байтом. В Python это одна строка из готовой библиотеки; в Go мы напишем парсер сами и заодно вспомним encoding/binary, compress/gzip и корректную обработку ошибок ввода-вывода.

Получаем файлы

Нужны четыре файла: картинки и метки для обучения и для теста.

train-images-idx3-ubyte.gz   # 60 000 картинок, ~9,9 МБ
train-labels-idx1-ubyte.gz   # 60 000 меток
t10k-images-idx3-ubyte.gz    # 10 000 картинок
t10k-labels-idx1-ubyte.gz    # 10 000 меток

Их раздают несколько зеркал (официальный сайт MNIST, зеркала на GitHub и в датасет-хабах). Скачай четыре .gz-файла в каталог data/mnist/ своего проекта — тот самый, который мы в первой главе добавили в .gitignore. Распаковывать вручную не нужно: читать gzip мы будем прямо из кода.

Формат IDX

IDX придуман для хранения многомерных массивов чисел и устроен предельно просто. Файл начинается с четырёхбайтового «магического числа»:

  • байты 0 и 1 — всегда нули;
  • байт 2 — код типа данных: 0x08 — беззнаковый байт, 0x09 — знаковый байт, 0x0B — int16, 0x0C — int32, 0x0D — float32, 0x0E — float64;
  • байт 3 — число измерений: 1 для меток, 3 для картинок.

Дальше идут размеры каждого измерения — по 4 байта на измерение, — а за ними подряд сами данные. Для картинок MNIST это 60000, 28, 28 и затем 47 040 000 байт яркости.

Критическая деталь: все числа записаны в порядке big-endian, то есть старшим байтом вперёд. Процессоры x86 и ARM работают в little-endian, поэтому читать надо через binary.BigEndian. Ошибка здесь даёт не падение, а бессмыслицу: вместо 60000 получишь 674 234 880, и make попытается выделить полтерабайта.

Отсюда и магические числа целиком: для картинок это 0x00000803 (2051 в десятичной), для меток — 0x00000801 (2049).

Читаем файл

package dataset

import (
    "compress/gzip"
    "encoding/binary"
    "fmt"
    "io"
    "os"
)

const (
    magicImages = 0x00000803 // 3 измерения, тип «беззнаковый байт»
    magicLabels = 0x00000801 // 1 измерение
)

// openMaybeGzip открывает файл и, если имя оканчивается на .gz,
// оборачивает его в gzip-ридер.
func openMaybeGzip(path string) (io.ReadCloser, error) {
    f, err := os.Open(path)
    if err != nil {
        return nil, fmt.Errorf("открыть %s: %w", path, err)
    }
    if filepath.Ext(path) != ".gz" {
        return f, nil
    }
    gz, err := gzip.NewReader(f)
    if err != nil {
        f.Close()
        return nil, fmt.Errorf("gzip %s: %w", path, err)
    }
    return gzipCloser{gz: gz, file: f}, nil
}

// gzipCloser закрывает и gzip-ридер, и нижележащий файл:
// иначе дескриптор утечёт.
type gzipCloser struct {
    gz   *gzip.Reader
    file *os.File
}

func (g gzipCloser) Read(p []byte) (int, error) { return g.gz.Read(p) }

func (g gzipCloser) Close() error {
    err := g.gz.Close()
    if cerr := g.file.Close(); err == nil {
        err = cerr
    }
    return err
}

Обёртка gzipCloser — тот случай, когда мелочь важна: gzip.Reader.Close() не закрывает исходный файл, и без этой обёртки при чтении четырёх файлов утекут четыре дескриптора. В долгоживущем сервисе такая утечка убивает процесс через сутки.

// ReadImages читает файл картинок IDX и возвращает
// срез изображений; каждое изображение — срез из rows*cols байт.
func ReadImages(path string) ([][]byte, int, int, error) {
    r, err := openMaybeGzip(path)
    if err != nil {
        return nil, 0, 0, err
    }
    defer r.Close()

    var header struct {
        Magic  int32
        Count  int32
        Rows   int32
        Cols   int32
    }
    // binary.Read с BigEndian заполнит все поля разом.
    if err := binary.Read(r, binary.BigEndian, &header); err != nil {
        return nil, 0, 0, fmt.Errorf("заголовок %s: %w", path, err)
    }
    if header.Magic != magicImages {
        return nil, 0, 0, fmt.Errorf("%s: magic %#x, ожидалось %#x", path, header.Magic, magicImages)
    }

    size := int(header.Rows * header.Cols)
    images := make([][]byte, header.Count)
    buf := make([]byte, size)
    for i := range images {
        // io.ReadFull обязателен: обычный Read вправе вернуть меньше байт,
        // чем просили, и наивный код тихо соберёт «сдвинутые» картинки.
        if _, err := io.ReadFull(r, buf); err != nil {
            return nil, 0, 0, fmt.Errorf("%s: картинка %d: %w", path, i, err)
        }
        images[i] = append([]byte(nil), buf...) // копия: buf переиспользуется
    }
    return images, int(header.Rows), int(header.Cols), nil
}

// ReadLabels читает файл меток IDX.
func ReadLabels(path string) ([]byte, error) {
    r, err := openMaybeGzip(path)
    if err != nil {
        return nil, err
    }
    defer r.Close()

    var header struct {
        Magic int32
        Count int32
    }
    if err := binary.Read(r, binary.BigEndian, &header); err != nil {
        return nil, fmt.Errorf("заголовок %s: %w", path, err)
    }
    if header.Magic != magicLabels {
        return nil, fmt.Errorf("%s: magic %#x, ожидалось %#x", path, header.Magic, magicLabels)
    }

    labels := make([]byte, header.Count)
    if _, err := io.ReadFull(r, labels); err != nil {
        return nil, fmt.Errorf("%s: метки: %w", path, err)
    }
    return labels, nil
}

Две вещи здесь стоят отдельного внимания. Первая — io.ReadFull: контракт io.Reader разрешает вернуть меньше данных, чем запрошено, и через gzip это происходит регулярно. Наивный r.Read(buf) в цикле по картинкам даст «сдвиг»: часть изображений склеится из хвоста одного и начала другого, а модель будет учиться на мусоре без единого сообщения об ошибке. Вторая — копия append([]byte(nil), buf...): без неё все элементы images будут ссылаться на один и тот же переиспользуемый буфер.

Подготовка данных для сети

Прочитанные байты в сеть подавать нельзя: нужны float64 в разумном диапазоне и метки в виде one-hot векторов.

// Dataset — готовый к обучению набор: X (n×784) и Y (n×10).
type Dataset struct {
    X *matrix.Matrix
    Y *matrix.Matrix
}

// Build превращает сырые байты в матрицы: пиксели нормируются в [0,1],
// метки разворачиваются в one-hot.
func Build(images [][]byte, labels []byte, pixels int) (*Dataset, error) {
    if len(images) != len(labels) {
        return nil, fmt.Errorf("картинок %d, меток %d", len(images), len(labels))
    }
    x := matrix.New(len(images), pixels)
    y := matrix.New(len(labels), 10)

    for i, img := range images {
        for j, p := range img {
            x.Data[i*pixels+j] = float64(p) / 255.0
        }
        lbl := labels[i]
        if lbl > 9 {
            return nil, fmt.Errorf("пример %d: метка %d вне диапазона 0..9", i, lbl)
        }
        y.Data[i*10+int(lbl)] = 1
    }
    return &Dataset{X: x, Y: y}, nil
}

Деление на 255 — та самая нормализация из первой главы. Без неё входы имеют размах 0…255, первый же слой выдаёт значения порядка сотен, ReLU их пропускает, следующий слой умножает ещё раз — и обучение либо расходится, либо требует микроскопического learning rate.

One-hot: цифра 3 превращается в вектор [0,0,0,1,0,0,0,0,0,0]. Почему не подать саму цифру одним числом? Потому что тогда модель решала бы задачу регрессии и считала, что 8 «ближе» к 9, чем к 1. Для классов такого порядка нет — восьмёрка не «больше» единицы, она просто другая.

Разбиение на обучение и валидацию

// Split делит датасет на две части в пропорции ratio,
// предварительно перемешав индексы детерминированно по сиду.
func Split(ds *Dataset, ratio float64, seed uint64) (train, val *Dataset) {
    rng := rand.New(rand.NewPCG(seed, seed+1))
    n := ds.X.Rows
    idx := make([]int, n)
    for i := range idx {
        idx[i] = i
    }
    rng.Shuffle(n, func(i, j int) { idx[i], idx[j] = idx[j], idx[i] })

    cut := int(float64(n) * ratio)
    return subset(ds, idx[:cut]), subset(ds, idx[cut:])
}

Валидационная выборка — данные, которые сеть не видит при обучении. Она отвечает на вопрос «модель научилась обобщать или просто запомнила примеры?». Перемешивание перед разрезом обязательно: MNIST не отсортирован по классам, но полагаться на это нельзя — во многих реальных датасетах данные сложены по группам, и наивный разрез «первые 80%» отдаёт в валидацию только последние классы.

Батч-итератор

// Batches возвращает функцию-итератор по батчам.
// Каждая эпоха начинается с нового перемешивания.
func (d *Dataset) Batches(size int, rng *rand.Rand) func() (*matrix.Matrix, *matrix.Matrix, bool) {
    idx := make([]int, d.X.Rows)
    for i := range idx {
        idx[i] = i
    }
    rng.Shuffle(len(idx), func(i, j int) { idx[i], idx[j] = idx[j], idx[i] })

    pos := 0
    return func() (*matrix.Matrix, *matrix.Matrix, bool) {
        if pos >= len(idx) {
            return nil, nil, false
        }
        end := min(pos+size, len(idx))
        bx := takeRows(d.X, idx[pos:end])
        by := takeRows(d.Y, idx[pos:end])
        pos = end
        return bx, by, true
    }
}

Замыкание вместо структуры-итератора — компактно и достаточно. Последний батч эпохи может оказаться неполным (60000 не делится на 64 нацело), и это нормально: функции потерь усредняют по фактическому числу строк.

Проверка глазами: рисуем цифру в терминале

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

// PrintASCII печатает изображение из строки матрицы X символами разной плотности.
func PrintASCII(ds *Dataset, row, cols int) {
    const shades = " .:-=+*#%@"
    for i := 0; i < ds.X.Cols/cols; i++ {
        line := make([]byte, cols)
        for j := 0; j < cols; j++ {
            v := ds.X.At(row, i*cols+j)        // яркость 0..1
            line[j] = shades[int(v*float64(len(shades)-1))]
        }
        fmt.Println(string(line))
    }
    fmt.Println("метка:", argmaxRow(ds.Y, row))
}

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

Кеширование: экономим время на каждом запуске

Разбор 60 000 картинок из gzip занимает пару секунд. За день экспериментов ты запустишь обучение сто раз, и эти секунды сложатся в неприятные минуты ожидания. Сохраним готовые матрицы в собственный бинарный кеш:

// SaveCache пишет матрицы в простой бинарный формат:
// [rows int32][cols int32][data float64...] для X, затем то же для Y.
func SaveCache(path string, ds *Dataset) error {
    f, err := os.Create(path)
    if err != nil {
        return fmt.Errorf("создать кеш: %w", err)
    }
    defer f.Close()

    w := bufio.NewWriter(f) // без буферизации запись миллионов float64 мучительно медленная
    for _, m := range []*matrix.Matrix{ds.X, ds.Y} {
        if err := binary.Write(w, binary.LittleEndian, int32(m.Rows)); err != nil {
            return err
        }
        if err := binary.Write(w, binary.LittleEndian, int32(m.Cols)); err != nil {
            return err
        }
        if err := binary.Write(w, binary.LittleEndian, m.Data); err != nil {
            return err
        }
    }
    return w.Flush() // забыть Flush — получить обрезанный файл
}

Здесь мы пишем в little-endian, потому что формат наш собственный и совместимость нужна только с самими собой. Обрати внимание на bufio.Writer: binary.Write без буферизации делает системный вызов на каждое значение, и запись 47 миллионов чисел растянется на минуты вместо секунды. И на Flush — без него хвост данных останется в буфере, а файл окажется битым.

Кейс: сборка полного пайплайна

// Load читает MNIST из каталога, при наличии кеша — из него.
func Load(dir string) (train, val, test *Dataset, err error) {
    cache := filepath.Join(dir, "mnist.cache")
    if ds, err := LoadCache(cache); err == nil {
        train, val = Split(ds, 0.9, 1)
        // тест хранится отдельно и никогда не участвует в обучении
    }
    // ...чтение IDX, Build, SaveCache при отсутствии кеша...
    return train, val, test, nil
}

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

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

1. Little-endian вместо big-endian

Число 60000 превращается в 674 234 880, программа пытается выделить сотни гигабайт и падает по нехватке памяти. Если make внезапно требует терабайт — первым делом проверяй порядок байтов.

2. Не проверять магическое число

Файл скачался наполовину, оказался HTML-страницей с ошибкой 404 или это вообще другой датасет — без проверки magic ты узнаешь об этом только по тому, что сеть не учится. Проверка стоит четыре байта и одну строку кода.

3. Read вместо io.ReadFull

Тихая порча данных: часть картинок собирается из кусков соседних. Не падает, не логируется, отлаживается неделю. Правило: если тебе нужно ровно N байт — io.ReadFull, всегда.

4. Забыть нормализацию

Сеть на входах 0…255 может обучиться при очень маленьком lr, но в разы медленнее и неустойчивее. Симптом: loss падает мучительно медленно или скачет при обычных значениях learning rate.

5. Утечка данных между выборками

Подбирать число нейронов по тестовой выборке — значит косвенно обучаться на ней. Тест открывают один раз, в конце. Если результат не понравился, честно говоря, надо возвращаться к валидации, а не «улучшать» по тесту.

6. Все картинки ссылаются на один буфер

images[i] = buf // все элементы указывают на одну и ту же память

В итоге весь датасет состоит из копий последней картинки. Классическая ловушка переиспользования буфера в Go — копируй явно.

Практика

Задание 1. Пакет dataset

Создай internal/dataset с ReadImages, ReadLabels, Build, Split, Batches. Все ошибки возвращай через fmt.Errorf с %w и указанием пути к файлу: сообщение «unexpected EOF» без имени файла бесполезно.

Задание 2. Тесты на маленьком синтетическом IDX

Сгенерируй в тесте крошечный IDX-файл в памяти (bytes.Buffer): magic, 2 картинки 3×3, затем 18 байт данных. Проверь, что парсер читает правильно; затем испорти magic и убедись, что возвращается ошибка. Так ты протестируешь загрузчик, не завися от скачанных файлов.

Задание 3. ASCII-визуализация

Реализуй PrintASCII и напечатай десять случайных примеров с метками. Убедись глазами, что цифры узнаваемы и метки совпадают. Это обязательный шаг: он один раз спасёт тебя от двух дней отладки «почему сеть не учится».

Задание 4. Статистика датасета

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

Задание 5. Кеш

Реализуй SaveCache и LoadCache с буферизацией. Замерь через time.Since время загрузки из IDX и из кеша, выведи оба числа. Убедись, что данные из кеша поэлементно совпадают с исходными.

Задание 6. Проверка разбиения

Напиши тест: после Split сумма размеров частей равна исходному размеру, множества индексов не пересекаются (добавь для этого временную версию, возвращающую индексы), а при одном сиде разбиение воспроизводится.

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

  • Загрузка сообщает: 60 000 обучающих и 10 000 тестовых примеров, размер 28×28.
  • Ошибочный magic и обрезанный файл дают понятные ошибки с именем файла.
  • ASCII-картинки узнаваемы, метки под ними совпадают с изображением.
  • Все значения X лежат в [0, 1], в каждой строке Y ровно одна единица.
  • Классы примерно сбалансированы: около 6000 примеров на цифру.
  • Загрузка из кеша заметно быстрее разбора IDX, данные идентичны.
  • Обучающая и валидационная выборки не пересекаются, тест не участвует нигде.

Итог

IDX — простой двоичный формат: магическое число, размеры измерений, данные подряд, всё в big-endian. Читать его надо через encoding/binary с binary.BigEndian и обязательно io.ReadFull, а gzip-обёртку закрывать вместе с файлом. Пиксели нормируются в [0, 1], метки разворачиваются в one-hot, данные делятся на обучение, валидацию и тест — и роли этих трёх выборок не смешиваются.

Данные готовы, сеть готова. В следующей главе мы соединим их: обучим полносвязный классификатор на MNIST, разберёмся с подбором learning rate, ранней остановкой и диагностикой по матрице ошибок.

Комментарии 0

Для добавления комментариев необходимо войти или зарегистрироваться.

Пока нет комментариев. Станьте первым!