Остановка градиента: tf.stopGradient

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

Основная концепция

tf.stopGradient принимает на вход тензор и возвращает новый тензор с теми же значениями, но градиенты для него не вычисляются. Это значит, что при вызове tf.grad или tf.variableGrads операции, зависящие от этого тензора, не будут влиять на обновление весов модели.

Синтаксис:

const y = tf.stopGradient(x);
  • x — исходный тензор.
  • y — тензор с заблокированными градиентами.

Пример использования

const x = tf.variable(tf.tensor1d([1, 2, 3]));
const y = x.square();

// Градиент будет вычислен
const dy_dx = tf.grad(x => x.square())(x);
dy_dx.print(); // [2, 4, 6]

// Градиент не будет распространяться
const z = tf.stopGradient(x).square();
const dz_dx = tf.grad(x => tf.stopGradient(x).square())(x);
dz_dx.print(); // [0, 0, 0]

В примере видно, что после применения tf.stopGradient функция dz_dx возвращает нули, поскольку градиент остановлен.

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

  1. Смешанные графы Иногда необходимо комбинировать части сети, где одна ветвь обучается, а другая используется только для вычислений без влияния на градиенты. tf.stopGradient позволяет исключить определённые операции из процесса обучения.

  2. Снижение расхода памяти Поскольку TensorFlow.js не сохраняет вычислительный граф для тензоров с остановленным градиентом, уменьшается количество промежуточных значений, необходимых для обратного прохода. Это особенно полезно при работе с большими моделями в браузере.

  3. Реализация target networks в reinforcement learning В алгоритмах вроде DQN используется копия сети (target network), градиенты которой не должны влиять на основную сеть. Использование tf.stopGradient гарантирует, что веса target network остаются неизменными при вычислении потерь.

Важные детали работы

  • Остановка градиента не меняет значения тензора, только блокирует обратное распространение.
  • Можно использовать tf.stopGradient внутри цепочек операций:
const a = tf.tensor1d([1, 2, 3]);
const b = a.mul(2);
const c = tf.stopGradient(b).add(1);
  • Градиент будет рассчитываться только для операций до tf.stopGradient. В примере градиент по a для операции c будет равен нулю.

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

При использовании оптимизаторов (tf.train.AdamOptimizer, tf.train.SGD) градиенты тензоров с остановленным градиентом не учитываются. Это обеспечивает корректное обучение только необходимых параметров:

const w = tf.variable(tf.randomNormal([2, 2]));
const x = tf.tensor2d([[1, 2]]);
const yTrue = tf.tensor2d([[0, 1]]);

const loss = () => {
  const logits = tf.stopGradient(x.matMul(w));
  return tf.losses.meanSquaredError(yTrue, logits);
};

const optimizer = tf.train.sgd(0.01);
optimizer.minimize(loss); // w не обновится

Лучшие практики

  • Применять tf.stopGradient только к тем тензорам, которые точно не должны участвовать в обновлении весов.
  • Для промежуточных вычислений без обратного распространения использовать локальные переменные (tf.tidy) совместно с остановкой градиента, чтобы снижать нагрузку на память.
  • В reinforcement learning и сложных архитектурах это инструмент контроля потока градиентов между разными ветвями сети.

Отличие от tf.variable и trainable: false

  • tf.stopGradient влияет только на обратное распространение, а не на сам тензор.
  • trainable: false при создании переменной блокирует её обновление оптимизатором, но не останавливает вычисление градиентов внутри операций с этой переменной. tf.stopGradient — более гибкий инструмент для блокировки градиентов выборочно.