Сохранение весов отдельно

ConvNetJS — это чисто JavaScript-библиотека для построения и обучения нейронных сетей, особенно сверточных (ConvNet), работающая полностью в браузере. Один из важных аспектов при работе с нейронными сетями — возможность сохранять веса отдельно от структуры сети, чтобы при необходимости загружать их в другую конфигурацию или использовать в различных проектах.

Архитектура хранения весов

В ConvNetJS веса слоёв хранятся в объектах Vol, которые представляют собой многомерные массивы данных с числовыми значениями параметров. Каждое соединение между нейронами и каждый фильтр сверточного слоя представлен массивом чисел. Внутри сети каждый слой имеет свои веса (weights) и смещения (biases), которые и подлежат сохранению.

Ключевые моменты:

  • layer.params — массив объектов параметров слоя.
  • Каждый параметр имеет поле w (веса) и dw (градиенты). Для сохранения интересуют только w.
  • Слой может содержать несколько Vol объектов, особенно это касается слоёв типа FullyConnected и Convolutional.

Сериализация весов

Для сохранения весов отдельно используется метод обхода всех слоёв и извлечения их параметров:

function saveWeights(net) {
    let weightsData = [];
    net.layers.forEach(layer => {
        if (layer.hasOwnProperty('filters')) { // сверточный слой
            layer.filters.forEach(filter => {
                weightsData.push(filter.w); 
            });
        }
        if (layer.hasOwnProperty('biases')) { // смещения
            weightsData.push(layer.biases.w);
        }
        if (layer.hasOwnProperty('params')) { // полностью связанные слои
            layer.params.forEach(p => {
                weightsData.push(p.w);
            });
        }
    });
    return JSON.stringify(weightsData);
}

Особенности:

  • Используется JSON.stringify, так как Vol представляет собой обычные массивы чисел.
  • Градиенты dw не сохраняются, так как для инференса они не нужны.
  • Слои без параметров игнорируются (например, ReLU или Pooling).

Загрузка весов

Для загрузки весов необходимо соответствие порядка параметров слоёв, так как ConvNetJS не хранит метаданные типа размеров фильтров внутри сериализованного массива. Пример функции восстановления:

function loadWeights(net, jsonWeights) {
    let weightsData = JSON.parse(jsonWeights);
    let idx = 0;
    net.layers.forEach(layer => {
        if (layer.hasOwnProperty('filters')) {
            layer.filters.forEach(filter => {
                filter.w = weightsData[idx++];
            });
        }
        if (layer.hasOwnProperty('biases')) {
            layer.biases.w = weightsData[idx++];
        }
        if (layer.hasOwnProperty('params')) {
            layer.params.forEach(p => {
                p.w = weightsData[idx++];
            });
        }
    });
}

Важные моменты:

  • Порядок загрузки должен строго соответствовать порядку сохранения.
  • Любое несоответствие размеров приведёт к ошибке или некорректной работе сети.
  • Если используется несколько разных сетей, хранить весовые файлы рекомендуется отдельно для каждой архитектуры.

Применение в реальных проектах

Сохранение весов отдельно позволяет:

  • Разделять обучение и инференс. Обучение может выполняться на сервере, а инференс — на клиенте.
  • Легко обмениваться обученными моделями без передачи всей структуры сети.
  • Сохранять промежуточные результаты для анализа прогресса обучения.

Советы по оптимизации

  • Для больших сетей и фильтров JSON может быть слишком тяжёлым. Возможна бинарная сериализация с использованием ArrayBuffer.
  • При многократной загрузке и сохранении полезно использовать контрольные суммы, чтобы убедиться, что данные не повреждены.
  • Для слоёв BatchNorm или других слоёв с дополнительными параметрами также сохраняются только w и biases.

Заключение по методике

Отделение весов от структуры сети в ConvNetJS обеспечивает гибкость в управлении моделями. Это позволяет создавать динамические системы, где одна и та же сеть может использовать разные наборы весов без изменения архитектуры, ускоряя эксперименты и упрощая деплой моделей в браузерных приложениях.