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

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

Основной принцип работы

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

const {value, grads} = tf.variableGrads(() => lossFn(), trainableVars);
  • value — значение функции потерь на текущем шаге.
  • grads — объект, где ключи соответствуют переменным, а значения — их градиенты.

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

Цикл обучения с tf.variableGrads

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

const learningRate = 0.01;
const x = tf.tensor1d([1, 2, 3, 4]);
const y = tf.tensor1d([2, 4, 6, 8]);

const w = tf.variable(tf.scalar(Math.random()));
const b = tf.variable(tf.scalar(Math.random()));

function predict(x) {
  return tf.add(tf.mul(x, w), b);
}

function loss(pred, label) {
  return pred.sub(label).square().mean();
}

for (let i = 0; i < 100; i++) {
  const {value, grads} = tf.variableGrads(() => loss(predict(x), y), [w, b]);

  w.assign(w.sub(grads[w].mul(learningRate)));
  b.assign(b.sub(grads[b].mul(learningRate)));

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

  tf.dispose([value, grads[w], grads[b]]);
}

Особенности подхода:

  • Ручное управление памятью: для предотвращения утечек нужно явно вызывать tf.dispose для градиентов и промежуточных значений.
  • Гибкость обновления: можно реализовать любые методы оптимизации, включая нестандартные формулы градиентного шага.
  • Поддержка нескольких переменных: tf.variableGrads возвращает объект с градиентами для всех переданных переменных, что удобно для сложных моделей с большим количеством параметров.

Совместимость с оптимизаторами

Даже при использовании tf.variableGrads можно сочетать ручное вычисление градиентов с оптимизаторами TensorFlow.js:

const optimizer = tf.train.sgd(learningRate);

for (let i = 0; i < 100; i++) {
  const {grads} = tf.variableGrads(() => loss(predict(x), y), [w, b]);
  optimizer.applyGradients(grads);
  tf.dispose(Object.values(grads));
}

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

Вычисление градиентов для сложных функций

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

function customLoss(x, y) {
  return tf.tidy(() => {
    let pred = x;
    for (let i = 0; i < 5; i++) {
      pred = pred.mul(w).add(b);
    }
    return pred.sub(y).square().mean();
  });
}

const {value, grads} = tf.variableGrads(() => customLoss(x, y), [w, b]);

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

Рекомендации по использованию

  • Всегда очищать память для градиентов и промежуточных значений.
  • Для больших моделей объединять tf.variableGrads с оптимизаторами для удобства и стабильности.
  • Для нестандартных обновлений параметров использовать градиенты напрямую.
  • При необходимости анализировать градиенты для диагностики проблемы переобучения или затухающего градиента.

Использование tf.variableGrads предоставляет полный контроль над процессом обучения в TensorFlow.js, позволяя реализовать как простые линейные модели, так и сложные нейронные сети с кастомными алгоритмами оптимизации.