Глава 10. Сохранение модели и инференс-сервис

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

Введение

Модель, живущая только в памяти процесса обучения, бесполезна. Финальная глава — про то, как довести её до продакшена: сохранить веса в файл, загрузить в другой программе, обернуть в HTTP-сервис, сделать его устойчивым к нагрузке и собрать в контейнер размером с сам бинарь.

Здесь пригодится всё, что ты знаешь про Go как про язык для сервисов: net/http, context, обработка сигналов, таймауты. Разница только в том, что за хендлером теперь стоит нейросеть — и у неё есть особенности, о которых обычный сервис не задумывается.

Формат весов

Первое решение — в чём хранить. Три варианта, и у каждого своя цена.

JSON. Читается глазами, отлаживается тривиально. Но 100 000 чисел в текстовом виде — это примерно 2 МБ вместо 800 КБ, парсинг заметно медленнее, а главное — float64 через текст округляется, и загруженная модель может давать чуть другие предсказания.

encoding/gob. Родной для Go бинарный формат, сам умеет структуры и срезы. Пишется в три строки. Минус — привязка к Go: прочитать веса из Python не выйдет. Для учебного проекта это чаще всего приемлемо.

Свой бинарный формат. Чуть больше кода, зато полный контроль: заголовок с магией и версией, компактность, читается из любого языка. Мы выберем его, потому что писать его учебно полезнее, а обратная совместимость получается явной.

package nn

const (
    modelMagic   uint32 = 0x474F4E4E // "GONN" в ASCII
    modelVersion uint32 = 1
)

// LayerSpec описывает слой в файле модели.
type LayerSpec struct {
    Kind string // "dense", "relu", "softmax"
    In   int32
    Out  int32
}

Магическое число в начале файла — та же идея, что в IDX из главы 6: мгновенно отличить свой файл от чужого и дать понятную ошибку вместо мусора. Версия формата обязательна: через полгода ты добавишь в слой ещё одно поле, и без версии старые файлы начнут читаться неправильно молча.

Сохранение архитектуры вместе с весами

Частая ошибка новичков — сохранить только числа. Тогда при загрузке нужно вручную собрать точно такую же сеть, и любое расхождение (128 вместо 256) даст либо панику, либо, что хуже, тихо неверный результат. Храни описание архитектуры в том же файле.

// Save записывает архитектуру и веса сети.
func Save(path string, net *Network) error {
    f, err := os.Create(path)
    if err != nil {
        return fmt.Errorf("создать %s: %w", path, err)
    }
    defer f.Close()

    w := bufio.NewWriter(f)
    if err := binary.Write(w, binary.LittleEndian, modelMagic); err != nil {
        return err
    }
    if err := binary.Write(w, binary.LittleEndian, modelVersion); err != nil {
        return err
    }

    specs := net.Specs() // описание слоёв
    if err := binary.Write(w, binary.LittleEndian, int32(len(specs))); err != nil {
        return err
    }
    for _, s := range specs {
        if err := writeSpec(w, s); err != nil {
            return err
        }
    }

    for _, l := range net.Layers {
        params, _ := l.Params()
        for _, p := range params {
            if err := binary.Write(w, binary.LittleEndian, p.Data); err != nil {
                return err
            }
        }
    }
    return w.Flush() // без Flush файл окажется обрезанным
}

// Load восстанавливает сеть из файла: сначала архитектуру, затем веса.
func Load(path string) (*Network, error) {
    f, err := os.Open(path)
    if err != nil {
        return nil, fmt.Errorf("открыть %s: %w", path, err)
    }
    defer f.Close()

    r := bufio.NewReader(f)
    var magic, version uint32
    if err := binary.Read(r, binary.LittleEndian, &magic); err != nil {
        return nil, fmt.Errorf("%s: заголовок: %w", path, err)
    }
    if magic != modelMagic {
        return nil, fmt.Errorf("%s: не файл модели (magic %#x)", path, magic)
    }
    if err := binary.Read(r, binary.LittleEndian, &version); err != nil {
        return nil, err
    }
    if version != modelVersion {
        return nil, fmt.Errorf("%s: версия формата %d, поддерживается %d", path, version, modelVersion)
    }
    // ...чтение спецификаций, сборка сети, заполнение весов...
    return net, nil
}

Детерминизм: проверяем, что сохранили то же самое

Единственный надёжный способ убедиться, что сериализация корректна, — сравнить предсказания до и после.

func TestSaveLoadRoundTrip(t *testing.T) {
    rng := rand.New(rand.NewPCG(1, 2))
    net := buildNet(rng)

    x := randomInput(rng, 8, 784)
    before := net.Forward(x)

    path := filepath.Join(t.TempDir(), "model.bin")
    if err := Save(path, net); err != nil {
        t.Fatalf("сохранение: %v", err)
    }
    loaded, err := Load(path)
    if err != nil {
        t.Fatalf("загрузка: %v", err)
    }

    after := loaded.Forward(x)
    for i := range before.Data {
        // Побитовое равенство: бинарный формат не теряет точность,
        // и любое расхождение означает ошибку в сериализации.
        if before.Data[i] != after.Data[i] {
            t.Fatalf("предсказание %d изменилось: %v против %v", i, before.Data[i], after.Data[i])
        }
    }
}

Здесь мы намеренно сравниваем через ==, хотя в главе 2 договорились так не делать. Причина: бинарная запись float64 сохраняет все биты, поэтому расхождение возможно только из-за ошибки — например, если ты записал float32 ради экономии. Для JSON-формата такой тест пришлось бы писать с допуском, и это ещё один аргумент против него.

Дополнительно полезно печатать при загрузке SHA-256 файла весов и писать его в лог сервиса — тогда по логу видно, какая именно модель обслуживает запросы.

HTTP-сервис распознавания

Теперь соберём cmd/serve. Сервис принимает изображение и возвращает распознанную цифру с вероятностями.

type predictRequest struct {
    // 784 значения яркости в диапазоне 0..1, построчно.
    Pixels []float64 `json:"pixels"`
}

type predictResponse struct {
    Digit  int       `json:"digit"`
    Probs  []float64 `json:"probs"`
    TookMs float64   `json:"took_ms"`
}

type server struct {
    pool  sync.Pool // пул сетей: см. раздел про потокобезопасность
    model string    // имя файла модели, для /healthz
}

func (s *server) handlePredict(w http.ResponseWriter, r *http.Request) {
    if r.Method != http.MethodPost {
        http.Error(w, "только POST", http.StatusMethodNotAllowed)
        return
    }

    // Ограничиваем размер тела: без этого один запрос может съесть память.
    r.Body = http.MaxBytesReader(w, r.Body, 1<<20)

    var req predictRequest
    if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
        http.Error(w, "некорректный JSON: "+err.Error(), http.StatusBadRequest)
        return
    }
    if len(req.Pixels) != 784 {
        http.Error(w, fmt.Sprintf("ожидалось 784 значения, получено %d", len(req.Pixels)),
            http.StatusBadRequest)
        return
    }

    start := time.Now()
    net := s.pool.Get().(*nn.Network)
    defer s.pool.Put(net)

    x := matrix.NewFrom(1, 784, req.Pixels)
    probs := nn.Softmax(net.Forward(x))

    resp := predictResponse{
        Digit:  argmaxRow(probs, 0),
        Probs:  probs.Data,
        TookMs: float64(time.Since(start).Microseconds()) / 1000,
    }

    w.Header().Set("Content-Type", "application/json")
    if err := json.NewEncoder(w).Encode(resp); err != nil {
        log.Printf("ответ клиенту: %v", err)
    }
}

Валидация входа здесь — не формальность. Сеть примет любой массив из 784 чисел, включая NaN и значения порядка 10⁹, и честно выдаст какой-то ответ. Проверять длину обязательно, а диапазон — крайне желательно: клиент, приславший 0…255 вместо 0…1, получит уверенные и совершенно неправильные предсказания.

Потокобезопасность: главная ловушка инференса

Помнишь lastInput из главы 5? Слой сохраняет вход при каждом Forward. HTTP-сервер обрабатывает запросы в отдельных горутинах — значит, два одновременных запроса будут писать в одно поле. Для чистого инференса результат Forward от этого не портится (кеш используется только в Backward), но полагаться на такое рассуждение опасно: добавишь dropout или batch norm — и получишь настоящую гонку. Детектор -race при этом ругается уже сейчас, потому что формально гонка есть.

Три честных решения:

  • Пул сетей (используем его выше): каждая горутина берёт свободный экземпляр. Веса можно разделять между копиями — они не меняются, — а кеши у каждой свои.
  • Мьютекс вокруг предсказания: просто, но сериализует все запросы и убивает пропускную способность.
  • Forward без побочных эффектов: отдельный метод инференса, ничего не кеширующий. Самое чистое решение, если инференс и обучение разнесены по коду.
// Пул готовых сетей: веса общие (они read-only), кеши — свои у каждой копии.
pool := sync.Pool{
    New: func() any {
        return net.CloneForInference() // копирует структуру, переиспользует матрицы весов
    },
}

Предобработка обязана совпадать с обучением

Самая обидная ошибка при выкатке модели: на обучении пиксели делились на 255, а сервис принимает 0…255 как есть. Модель формально работает, точность — на уровне случайного угадывания, и никаких сообщений об ошибке. Ровно то же с инверсией цвета: MNIST — белые цифры на чёрном фоне, а картинка от пользователя обычно чёрная на белом, и её нужно инвертировать.

// decodePNG приводит присланную картинку к формату обучающих данных:
// оттенки серого, 28×28, белая цифра на чёрном фоне, значения 0..1.
func decodePNG(r io.Reader) ([]float64, error) {
    img, err := png.Decode(r)
    if err != nil {
        return nil, fmt.Errorf("png: %w", err)
    }
    b := img.Bounds()
    if b.Dx() != 28 || b.Dy() != 28 {
        return nil, fmt.Errorf("ожидалось 28x28, получено %dx%d", b.Dx(), b.Dy())
    }

    out := make([]float64, 28*28)
    for y := 0; y < 28; y++ {
        for x := 0; x < 28; x++ {
            gray, _, _, _ := img.At(b.Min.X+x, b.Min.Y+y).RGBA()
            v := float64(gray) / 65535.0 // RGBA() возвращает 16-битные значения
            out[y*28+x] = 1 - v          // инверсия: фон должен быть нулём
        }
    }
    return out, nil
}

Практическое правило: код предобработки должен быть один и тот же для обучения и для инференса. Лучше всего вынести его в общую функцию пакета dataset и вызывать из обоих мест — тогда рассинхронизация физически невозможна.

Таймауты и graceful shutdown

func main() {
    modelPath := flag.String("model", "model.bin", "файл весов")
    addr := flag.String("addr", ":8080", "адрес сервиса")
    flag.Parse()

    net, err := nn.Load(*modelPath)
    if err != nil {
        log.Fatalf("модель: %v", err)
    }
    log.Printf("модель загружена: %s", *modelPath)

    srv := &http.Server{
        Addr:    *addr,
        Handler: newRouter(net),
        // Таймауты обязательны: сервер без них уязвим к медленным клиентам,
        // которые держат соединения открытыми и исчерпывают дескрипторы.
        ReadTimeout:       5 * time.Second,
        ReadHeaderTimeout: 2 * time.Second,
        WriteTimeout:      10 * time.Second,
        IdleTimeout:       60 * time.Second,
    }

    // Контекст, отменяемый по SIGINT/SIGTERM.
    ctx, stop := signal.NotifyContext(context.Background(),
        os.Interrupt, syscall.SIGTERM)
    defer stop()

    go func() {
        if err := srv.ListenAndServe(); err != nil && !errors.Is(err, http.ErrServerClosed) {
            log.Fatalf("сервер: %v", err)
        }
    }()
    log.Printf("слушаю %s", *addr)

    <-ctx.Done()
    log.Println("останавливаюсь...")

    // Даём текущим запросам доработать, новые не принимаем.
    shutdownCtx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
    defer cancel()
    if err := srv.Shutdown(shutdownCtx); err != nil {
        log.Printf("принудительное завершение: %v", err)
    }
}

Без graceful shutdown при деплое обрываются запросы, которые уже в обработке. С ним контейнер получает SIGTERM, перестаёт принимать новые соединения, дорабатывает текущие и завершается — пользователи ничего не замечают.

Health-check

// /healthz — жив ли процесс; /readyz — готов ли обслуживать (модель загружена).
mux.HandleFunc("/healthz", func(w http.ResponseWriter, r *http.Request) {
    w.WriteHeader(http.StatusOK)
    fmt.Fprintln(w, "ok")
})

Два разных эндпоинта нужны потому, что оркестратор задаёт два разных вопроса: «не завис ли процесс, не перезапустить ли его» и «можно ли уже слать сюда трафик». Модель грузится секунду-другую, и в это время сервис жив, но не готов.

Сборка: один бинарь и пустой образ

CGO_ENABLED=0 go build -ldflags="-s -w" -o bin/serve ./cmd/serve

CGO_ENABLED=0 даёт полностью статический бинарь без зависимости от системного libc — именно это позволяет положить его в пустой образ. Флаги -s -w выбрасывают отладочные таблицы и уменьшают размер примерно на треть.

FROM golang:1.23 AS build
WORKDIR /src
COPY go.mod ./
RUN go mod download
COPY . .
RUN CGO_ENABLED=0 go build -ldflags="-s -w" -o /out/serve ./cmd/serve

FROM scratch
COPY --from=build /out/serve /serve
COPY model.bin /model.bin
EXPOSE 8080
ENTRYPOINT ["/serve", "-model", "/model.bin", "-addr", ":8080"]

Итоговый образ — это бинарь плюс файл весов, единицы мегабайт. Ни интерпретатора, ни пакетного менеджера, ни системных библиотек: и размер меньше, и поверхность атаки уже. Если сервису понадобится HTTPS к внешним хостам, вместо scratch берут gcr.io/distroless/static — там есть корневые сертификаты.

Кейс: замер задержки

func BenchmarkPredict(b *testing.B) {
    net, err := nn.Load("testdata/model.bin")
    if err != nil {
        b.Fatal(err)
    }
    x := matrix.New(1, 784)

    b.ResetTimer()
    b.ReportAllocs()
    for i := 0; i < b.N; i++ {
        _ = nn.Softmax(net.Forward(x))
    }
}

// Параллельная версия: сколько запросов в секунду выдержит сервис.
func BenchmarkPredictParallel(b *testing.B) {
    net, _ := nn.Load("testdata/model.bin")
    pool := sync.Pool{New: func() any { return net.CloneForInference() }}

    b.RunParallel(func(pb *testing.PB) {
        x := matrix.New(1, 784)
        for pb.Next() {
            n := pool.Get().(*nn.Network)
            _ = nn.Softmax(n.Forward(x))
            pool.Put(n)
        }
    })
}

Ориентир для сети 784-128-10: одно предсказание — порядка 50–150 микросекунд, то есть тысячи запросов в секунду с одного ядра. Узким местом при таких числах становится не модель, а JSON-сериализация и сеть — что хорошо иллюстрирует, почему профилировать надо весь путь запроса, а не только матричную арифметику.

Проверить сервис вручную:

go run ./cmd/serve -model model.bin &
curl -s -X POST localhost:8080/predict \
  -H 'Content-Type: application/json' \
  -d "{\"pixels\": $(python3 -c 'import json;print(json.dumps([0.0]*784))')}" | head

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

1. Сохранить веса без архитектуры

Через месяц ты не вспомнишь, было там 128 нейронов или 256, и файл превратится в набор чисел неизвестного назначения. Пиши спецификацию слоёв в тот же файл.

2. Формат без версии и магии

Любое изменение структуры молча ломает старые файлы, а чужой файл читается как мусор. Четыре байта магии и четыре байта версии стоят дёшево и снимают целый класс проблем.

3. Разная предобработка при обучении и инференсе

Забытая нормализация или неинвертированный фон дают точность на уровне случайного угадывания при полностью «рабочем» сервисе. Общая функция предобработки — единственная надёжная защита.

4. Сервер без таймаутов

http.ListenAndServe(addr, handler) в одну строку удобно для примеров и опасно в проде: медленный клиент удерживает соединение сколь угодно долго. Всегда создавай http.Server явно и задавай таймауты.

5. Тело запроса без ограничения размера

Без http.MaxBytesReader клиент может прислать гигабайт JSON и исчерпать память процесса. Одна строка защиты.

6. Общая сеть на все запросы

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

7. Забыть Flush у bufio.Writer

Файл модели окажется обрезанным, а ошибка проявится только при загрузке — возможно, через неделю. return w.Flush() последней строкой, всегда.

Практика

Задание 1. Формат модели

Реализуй Save и Load со своим бинарным форматом: магия, версия, спецификации слоёв, веса. Обработай три ошибки явными сообщениями: неверная магия, неподдерживаемая версия, обрезанный файл.

Задание 2. Тест round-trip

Напиши тест «обучил → сохранил → загрузил → предсказания побитово совпали». Затем испорти файл (обрежь последние 100 байт) и убедись, что Load возвращает понятную ошибку, а не панику.

Задание 3. HTTP-сервис

Собери cmd/serve с эндпоинтами POST /predict (JSON с массивом пикселей), POST /predict/png (изображение 28×28), GET /healthz и GET /readyz. Обязательно: таймауты сервера, MaxBytesReader, валидация длины входа, JSON-ответ с кодами ошибок.

Задание 4. Потокобезопасность

Напиши тест, который шлёт 100 одновременных запросов через httptest.NewServer и проверяет, что все ответы одинаковы для одного и того же входа. Прогони с -race. Затем убери пул сетей и посмотри, что скажет детектор.

Задание 5. Graceful shutdown

Добавь остановку по сигналу через signal.NotifyContext и srv.Shutdown. Проверь вручную: запусти сервис, начни долгий запрос, нажми Ctrl+C — запрос должен завершиться, а процесс выйти без ошибок.

Задание 6. Контейнер

Напиши многоступенчатый Dockerfile со сборкой на golang и финальным образом на scratch. Собери, запусти, проверь curl. Сравни размер образа с образом на базе python:3.12 — разница обычно в 50–100 раз.

Задание 7. Бенчмарк задержки

Замерь время одного предсказания и пропускную способность через b.RunParallel. Отдельно замерь полный путь через HTTP (httptest) и посчитай, какую долю времени занимает сама сеть, а какую — JSON и сетевой стек.

Задание 8. Свои цифры

Нарисуй в любом редакторе несколько цифр 28×28, сохрани в PNG и отправь сервису. Скорее всего, точность окажется хуже, чем на MNIST, — и это отличный финальный урок: модель работает ровно на тех данных, которые похожи на обучающие. Попробуй понять, чем твои картинки отличаются (толщина линии, центрирование, контраст), и приведи их к формату MNIST.

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

  • Загрузка чужого или битого файла даёт понятную ошибку, а не панику.
  • Предсказания до и после сохранения совпадают побитово.
  • Сервис отвечает на /predict корректным JSON и возвращает 400 на некорректный вход.
  • go test -race ./... проходит чисто, включая тест с 100 параллельными запросами.
  • У http.Server заданы все четыре таймаута, тело запроса ограничено.
  • Ctrl+C останавливает сервис без обрыва текущих запросов.
  • Образ на scratch собирается и работает; в нём только бинарь и веса.
  • Предобработка на инференсе — та же функция, что и при обучении.

Итог

Модель сохраняется вместе с описанием архитектуры, в бинарном формате с магическим числом и версией; корректность сериализации проверяется тестом на побитовое совпадение предсказаний. Инференс-сервис — обычный Go-сервис с явными таймаутами, ограничением размера тела, health-check'ами и graceful shutdown, но с одной особенностью: кеши слоёв делают сеть непотокобезопасной, и одновременные запросы требуют пула экземпляров либо инференса без побочных эффектов. Предобработка на сервисе обязана быть той же самой функцией, что при обучении.

Что дальше

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

Куда двигаться, если хочется продолжения: реализовать batch normalization и полноценный обратный проход свёртки; попробовать другие датасеты (Fashion-MNIST — та же загрузка, задача заметно сложнее); добавить аугментацию данных (сдвиги и повороты картинок) и посмотреть, как вырастет устойчивость; собрать рекуррентный слой и поработать с последовательностями. И, конечно, взять любой известный фреймворк — теперь его исходники читаются как знакомый текст.

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

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

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