Отладка градиентов

TensorFlow.js предоставляет высокоуровневый API для построения и обучения моделей машинного обучения в JavaScript. Одной из ключевых концепций является автоматическое дифференцирование, позволяющее вычислять градиенты функций потерь относительно параметров модели. Градиенты необходимы для оптимизации весов с помощью алгоритмов, таких как стохастический градиентный спуск (SGD).

Функция tf.variable и обновление параметров

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

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

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

Автоматическое вычисление градиентов

TensorFlow.js использует функцию tf.grads или tf.variableGrads для получения градиентов. Более гибкий инструмент — tf.tidy с tf.variableGrads для предотвращения утечек памяти.

Пример вычисления градиентов для простой линейной модели:

const f = (w, b) => tf.tidy(() => {
  const x = tf.tensor1d([1, 2, 3, 4]);
  const y = tf.tensor1d([1, 3, 5, 7]);
  const yPred = x.mul(w).add(b);
  return yPred.sub(y).square().mean();
});

const {value, grads} = tf.variableGrads(() => f(w, b));
console.log(value.dataSync());
console.log(grads[w].dataSync());
console.log(grads[b].dataSync());

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

Отслеживание операций с tf.GradientTape

Начиная с последних версий TensorFlow.js, появилась возможность использовать концепцию “ленты градиентов”, аналогичную TensorFlow Python:

const x = tf.variable(tf.scalar(2));
const y = tf.variable(tf.scalar(3));

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

tf.tidy(() => {
  const f = () => x.square().add(y.square());
  
  const grads = tf.grads(f);
  const [dx, dy] = grads([x, y]);

  x.assign(x.sub(dx.mul(0.1)));
  y.assign(y.sub(dy.mul(0.1)));
});

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

tf.tidy(() => {
  tf.variableGrads(() => {
    const tape = tf.grad((x) => x.square().add(10));
    const grad = tape(x);
  });
});

Проверка градиентов

Отладка градиентов включает несколько аспектов:

  1. Визуальная проверка численных значений: градиенты должны иметь разумный масштаб и знак.
  2. Проверка конечной производной: для функции (f(x) = x^2), градиент должен быть (2x).
  3. Использование tf.print: для небольших тензоров можно выводить значения градиентов на каждом шаге:
tf.print(grads[w]);
tf.print(grads[b]);
  1. Сравнение с численным градиентом: иногда полезно вычислить производную через конечные разности и убедиться, что результат совпадает с автоматическим градиентом.

Особенности работы с графом вычислений

  • tf.tidy освобождает память автоматически. Все промежуточные тензоры, не используемые за пределами функции, будут удалены. Это особенно важно при обучении больших моделей.
  • Градиенты по нескольким переменным: TensorFlow.js позволяет одновременно вычислять градиенты по всем переменным модели, что упрощает обучение сложных нейронных сетей.
  • Изоляция вычислений: если градиенты вычисляются внутри вложенных функций, нужно убедиться, что не происходит утечек памяти, а переменные остаются доступными для оптимизатора.

Практическое применение

  • Настройка скорости обучения и оптимизатора напрямую влияет на величину и поведение градиентов.
  • Взрыв градиентов и исчезновение градиентов особенно актуальны для рекуррентных сетей. TensorFlow.js поддерживает операции clipByValue и clipByNorm для стабилизации градиентов:
const clippedGrad = grads[w].clipByValue(-1, 1);
  • Визуализация градиентов через графики или консоль помогает определить неправильные вычисления или ошибки в архитектуре модели.

Встроенные функции для анализа

TensorFlow.js предоставляет дополнительные инструменты для отладки:

  • tf.nextFrame() — позволяет визуализировать изменения в веб-браузере без блокировки интерфейса.
  • tf.memory() — показывает количество занятых тензоров, полезно для отслеживания утечек памяти при частых вычислениях градиентов.

Эти методы в совокупности создают полноценную систему для эффективной и безопасной отладки градиентов в браузере или Node.js.