В 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 возвращает нули, поскольку градиент
остановлен.
Смешанные графы Иногда необходимо комбинировать
части сети, где одна ветвь обучается, а другая используется только для
вычислений без влияния на градиенты. tf.stopGradient
позволяет исключить определённые операции из процесса обучения.
Снижение расхода памяти Поскольку TensorFlow.js не сохраняет вычислительный граф для тензоров с остановленным градиентом, уменьшается количество промежуточных значений, необходимых для обратного прохода. Это особенно полезно при работе с большими моделями в браузере.
Реализация 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) совместно с
остановкой градиента, чтобы снижать нагрузку на память.tf.variable и trainable: falsetf.stopGradient влияет только на обратное
распространение, а не на сам тензор.trainable: false при создании переменной блокирует её
обновление оптимизатором, но не останавливает вычисление градиентов
внутри операций с этой переменной. tf.stopGradient — более
гибкий инструмент для блокировки градиентов выборочно.