Ручное обновление весов

В TensorFlow.js обучение модели обычно автоматизировано с помощью методов оптимизации, таких как optimizer.minimize. Однако для глубокого понимания работы нейронных сетей важно освоить ручное управление весами и градиентами. Это позволяет более точно контролировать процесс обучения, реализовать нестандартные алгоритмы обновления и эксперименты с нестандартными функциями потерь.

Создание переменных для весов

В TensorFlow.js параметры модели представлены объектами tf.Variable. Они хранят значения, которые можно изменять во время обучения.

const w = tf.variable(tf.randomNormal([3, 3])); // матрица весов 3x3
const b = tf.variable(tf.zeros([3]));          // вектор смещений
  • tf.variable создаёт переменную, которую можно изменять в процессе вычислений.
  • tf.randomNormal и tf.zeros используются для начальной инициализации весов.

Вычисление градиентов вручную

Для ручного обновления весов необходимо получить градиенты функции потерь относительно этих переменных. TensorFlow.js предоставляет функцию tf.grads и tf.variableGrads для этой задачи.

const x = tf.tensor2d([[1, 2, 3]]);

function lossFunction() {
    const yPred = x.matMul(w).add(b);
    const yTrue = tf.tensor2d([[0.5, 1, 1.5]]);
    return yPred.sub(yTrue).square().mean();
}

const grads = tf.variableGrads(lossFunction);
  • grads.value — значение функции потерь.
  • grads.grads — объект, содержащий градиенты по каждой переменной.

Прямое обновление весов

После получения градиентов можно обновлять веса вручную с использованием выбранного темпа обучения.

const learningRate = 0.01;

w.assign(w.sub(grads.grads[w].mul(learningRate)));
b.assign(b.sub(grads.grads[b].mul(learningRate)));
  • assign заменяет текущее значение переменной новым.
  • sub и mul выполняют арифметические операции над тензорами.

Ручное обновление позволяет экспериментировать с различными правилами корректировки весов, включая адаптивные схемы, которые невозможно реализовать через стандартные оптимизаторы.

Использование циклов обучения

Для многократного обновления весов используется цикл обучения:

for (let i = 0; i < 1000; i++) {
    const {value, grads} = tf.variableGrads(lossFunction);
    w.assign(w.sub(grads[w].mul(learningRate)));
    b.assign(b.sub(grads[b].mul(learningRate)));

    if (i % 100 === 0) {
        console.log(`Step ${i}: loss = ${value.dataSync()}`);
    }

    tf.dispose([value, grads[w], grads[b]]);
}
  • Цикл позволяет контролировать каждый шаг обучения.
  • tf.dispose освобождает память от промежуточных тензоров, предотвращая утечки памяти.

Применение сложных схем обновления

Ручное управление весами открывает возможности для:

  • Внедрения динамического темпа обучения, который изменяется по шагам или эпохам.
  • Реализации алгоритмов с моментумом без стандартного оптимизатора.

Пример с моментумом:

let velocity = tf.zerosLike(w);
const momentum = 0.9;

for (let i = 0; i < 500; i++) {
    const {grads} = tf.variableGrads(lossFunction);
    velocity = velocity.mul(momentum).sub(grads[w].mul(learningRate));
    w.assign(w.add(velocity));

    tf.dispose([grads[w], grads[b]]);
}
  • velocity аккумулирует предыдущее обновление для сглаживания градиентов.
  • Такой подход позволяет повысить стабильность обучения и ускорить сходимость.

Рекомендации по управлению памятью

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

  • Использовать tf.tidy для автоматического удаления временных тензоров:
tf.tidy(() => {
    const {grads} = tf.variableGrads(lossFunction);
    w.assign(w.sub(grads[w].mul(learningRate)));
});
  • Освобождать явным вызовом dispose тензоры, которые сохраняются вне tf.tidy.

Вывод

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