Контроль обновления весов вручную

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

Основы градиентного обновления

Обновление весов нейронной сети — это процесс корректировки параметров модели на основе вычисленных градиентов функции потерь. Формально, для веса ( w ) и функции потерь ( L ) шаг обновления с использованием градиентного спуска выглядит так:

[ w := w - ]

где ( ) — коэффициент обучения (learning rate), а ( ) — градиент функции потерь по весу ( w ).

В TensorFlow.js градиенты вычисляются с помощью функции tf.grad или tf.variableGrads, что позволяет использовать их для ручного обновления параметров модели.

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

Для ручного управления весами все параметры модели должны быть объявлены как tf.Variable. Пример:

const w1 = tf.variable(tf.randomNormal([inputSize, hiddenSize]));
const b1 = tf.variable(tf.zeros([hiddenSize]));
const w2 = tf.variable(tf.randomNormal([hiddenSize, outputSize]));
const b2 = tf.variable(tf.zeros([outputSize]));

Использование tf.variable позволяет изменять значения тензоров после их создания, что необходимо для прямого применения градиентного спуска вручную.

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

Для вычисления градиентов используется функция tf.variableGrads, которая принимает функцию потерь и возвращает объект с градиентами для всех переменных:

const loss = (pred, label) => pred.sub(label).square().mean();

const computeGradients = (inputs, labels) => {
  return tf.variableGrads(() => {
    const hidden = inputs.matMul(w1).add(b1).relu();
    const output = hidden.matMul(w2).add(b2);
    return loss(output, labels);
  });
};

В объекте, возвращаемом tf.variableGrads, ключи — это переменные, а значения — соответствующие градиенты.

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

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

const learningRate = 0.01;

const applyGradients = (grads) => {
  w1.assign(w1.sub(grads[w1].mul(learningRate)));
  b1.assign(b1.sub(grads[b1].mul(learningRate)));
  w2.assign(w2.sub(grads[w2].mul(learningRate)));
  b2.assign(b2.sub(grads[b2].mul(learningRate)));
};

assign заменяет текущее значение переменной новым тензором. Выражение w1.sub(grads[w1].mul(learningRate)) создаёт обновлённый тензор для веса с учётом шага обучения.

Использование tf.tidy для управления памятью

При ручном обновлении весов важно контролировать использование памяти, так как TensorFlow.js не освобождает автоматически промежуточные тензоры. Для этого все вычисления можно помещать в tf.tidy:

tf.tidy(() => {
  const grads = computeGradients(inputs, labels);
  applyGradients(grads.grads);
});

tf.tidy гарантирует удаление всех промежуточных тензоров, которые больше не нужны, предотвращая утечки памяти в браузере.

Применение кастомного алгоритма оптимизации

Ручное управление весами позволяет реализовать любой алгоритм оптимизации. Например, стохастический градиентный спуск с импульсом:

let vW1 = tf.zerosLike(w1);
let vB1 = tf.zerosLike(b1);
let vW2 = tf.zerosLike(w2);
let vB2 = tf.zerosLike(b2);
const momentum = 0.9;

const applyMomentum = (grads) => {
  vW1 = vW1.mul(momentum).sub(grads[w1].mul(learningRate));
  vB1 = vB1.mul(momentum).sub(grads[b1].mul(learningRate));
  vW2 = vW2.mul(momentum).sub(grads[w2].mul(learningRate));
  vB2 = vB2.mul(momentum).sub(grads[b2].mul(learningRate));

  w1.assign(w1.add(vW1));
  b1.assign(b1.add(vB1));
  w2.assign(w2.add(vW2));
  b2.assign(b2.add(vB2));
};

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

Интеграция с обучающим циклом

Полный цикл обучения с ручным обновлением весов выглядит следующим образом:

for (let epoch = 0; epoch < epochs; epoch++) {
  tf.tidy(() => {
    const grads = computeGradients(inputs, labels);
    applyGradients(grads.grads);
  });
  if (epoch % 10 === 0) {
    const currentLoss = loss(forward(inputs), labels).dataSync()[0];
    console.log(`Epoch ${epoch}: Loss = ${currentLoss}`);
  }
}

forward(inputs) — функция прямого прохода, аналогичная вычислению предсказания в computeGradients. Этот цикл позволяет полностью контролировать процесс обновления весов на каждом шаге.

Преимущества ручного обновления весов

  • Возможность экспериментировать с нестандартными алгоритмами оптимизации.
  • Контроль над распределением градиентов и скоростью обучения.
  • Гибкость для реализации кастомных механизмов регуляризации.
  • Полный доступ к промежуточным вычислениям для диагностики и отладки.

Использование этих подходов в TensorFlow.js открывает путь к глубокому пониманию работы нейронных сетей и расширяет возможности экспериментов с обучением моделей на стороне клиента.