Прунинг (pruning) и сжатие моделей являются важными методами оптимизации нейронных сетей для уменьшения их размера, ускорения инференса и снижения потребления памяти. В контексте TensorFlow.js эти методы особенно актуальны для веб-приложений, где ресурсы ограничены и время загрузки модели критично.
Прунинг заключается в удалении или обнулении наименее значимых весов нейронной сети. Основная идея — сократить количество вычислений без значительной потери точности модели.
Типы прунинга:
Критерии выбора весов для обнуления:
В TensorFlow.js можно реализовать прунинг с помощью перебора весов
модели через API tf.layers или tf.Model. Для
каждого слоя извлекаются веса методом getWeights(), затем
выполняется маскирование весов с минимальными значениями:
const pruneWeights = (weights, threshold) => {
return weights.map(w => {
const data = w.arraySync();
const prunedData = data.map(v => Math.abs(v) < threshold ? 0 : v);
return tf.tensor(prunedData, w.shape);
});
};
const layer = model.layers[0];
const prunedWeights = pruneWeights(layer.getWeights(), 0.01);
layer.setWeights(prunedWeights);
Квантование — процесс преобразования весов и активаций модели из формата с плавающей запятой (float32) в более компактные форматы (int8, float16). Квантование уменьшает размер модели и ускоряет вычисления на устройствах с ограниченными ресурсами.
Методы квантования:
В TensorFlow.js можно использовать утилиту
tfjs-converter для пост-тренировочного квантования при
конвертации моделей из формата TensorFlow SavedModel:
tensorflowjs_converter \
--input_format=tf_saved_model \
--quantization_bytes=1 \
/saved_model_path \
/tfjs_model_path
Сжатие может касаться не только весов, но и самой архитектуры сети:
Для динамического прунинга важно правильно сохранять маски, чтобы при дальнейшем обучении обнуленные веса не восстанавливались случайно. В TensorFlow.js маски можно хранить в виде бинарных тензоров и применять их на каждом шаге обновления весов:
const applyMask = (weights, mask) => {
return weights.mul(mask);
};
После прунинга и сжатия модель можно сохранить для веб-приложений:
await model.save('downloads://pruned_model');
Это создает .json и бинарные .bin файлы,
оптимизированные для загрузки и использования в браузере.
Эти методы обеспечивают баланс между компактностью модели и сохранением качества предсказаний, что особенно важно для интерактивных веб-приложений, где каждый килобайт и миллисекунда имеют значение.