Обработка NaN и Inf в потере

В машинном обучении крайне важно контролировать корректность значений, возникающих на различных этапах вычислений. В JavaScript с TensorFlow.js это особенно актуально, так как библиотека выполняет операции с числами с плавающей запятой, где могут появляться значения NaN (Not a Number) и Infinity (Inf). Эти значения нарушают процесс обучения и могут приводить к остановке градиентного спуска или некорректным обновлениям весов.


Причины появления NaN и Inf

  1. Деление на ноль Любая операция деления, где знаменатель равен нулю, приводит к Infinity или NaN при 0/0. Пример:

    const a = tf.scalar(0);
    const b = tf.scalar(0);
    const c = a.div(b); // NaN
  2. Логарифм нуля или отрицательного числа Использование tf.log(x) при x <= 0 вызовет -Infinity или NaN.

    const x = tf.scalar(0);
    const y = tf.log(x); // -Infinity
  3. Большие значения и переполнение Экспоненциальные функции (tf.exp) могут выдавать Infinity при слишком больших входных значениях.

    const x = tf.scalar(1000);
    const y = tf.exp(x); // Infinity
  4. Нестабильность численных методов Градиентный спуск с слишком большим learning rate может привести к резкому росту весов и NaN в функции потерь.


Обнаружение NaN и Inf

TensorFlow.js предоставляет утилиты для проверки числовых значений:

const tensor = tf.tensor([1, 2, NaN, Infinity, 5]);

const hasNaN = tensor.isNaN().any().dataSync()[0];      // true
const hasInf = tensor.isInf().any().dataSync()[0];      // true
  • isNaN() возвращает тензор булевых значений, где true соответствует NaN.
  • isInf() аналогично проверяет бесконечные значения.

Использование any() и dataSync() позволяет получить единичный флаг наличия проблемных значений.


Обработка NaN и Inf в функции потерь

Замена значений

Для безопасного обучения необходимо предотвращать появление NaN и Inf. Основной метод — замена проблемных значений на допустимые:

function safeLoss(yTrue, yPred) {
  let loss = tf.losses.meanSquaredError(yTrue, yPred);

  loss = tf.where(
    loss.isNaN().logicalOr(loss.isInf()),
    tf.zerosLike(loss),   // Заменяем на 0
    loss
  );

  return loss;
}
  • tf.where(condition, x, y) позволяет выбрать между значениями x и y по условию condition.
  • zerosLike создаёт тензор той же формы, что и loss, заполненный нулями.

Клиппинг градиентов

Для предотвращения NaN и Inf при обучении можно ограничивать значения градиентов:

const optimizer = tf.train.adam(0.01);

optimizer.minimize(() => {
  const preds = model.predict(inputs);
  const loss = tf.losses.meanSquaredError(targets, preds);

  return tf.clipByValue(loss, -1e5, 1e5); // Ограничение потерь
});
  • tf.clipByValue(tensor, min, max) безопасно ограничивает значения тензора заданными пределами.
  • Это предотвращает экстремальные значения, которые могут вызвать переполнение.

Предотвращение NaN при логарифмах и делениях

function logSafe(x) {
  const epsilon = 1e-7; // Малое число для стабильности
  return tf.log(x.add(tf.scalar(epsilon)));
}

function divSafe(a, b) {
  const epsilon = 1e-7;
  return a.div(b.add(tf.scalar(epsilon)));
}
  • Добавление epsilon гарантирует, что знаменатель не равен нулю, а логарифм не вычисляется от нуля.
  • Этот метод применяется в большинстве функций потерь с логарифмами (cross-entropy) и нормализованных делениях.

Мониторинг потерь во время обучения

Для раннего выявления NaN и Inf полезно отслеживать значения функции потерь на каждом шаге:

for (let i = 0; i < epochs; i++) {
  optimizer.minimize(() => {
    const preds = model.predict(inputs);
    const loss = tf.losses.meanSquaredError(targets, preds);

    const lossValue = loss.dataSync()[0];
    if (!isFinite(lossValue)) {
      console.warn(`Проблемное значение потерь на эпохе ${i}: ${lossValue}`);
    }

    return loss;
  });
}
  • isFinite(value) проверяет, является ли число конечным.
  • Предупреждения помогают выявить нестабильность модели на раннем этапе.

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

  1. Всегда проверять входные данные на NaN и Inf перед обучением.
  2. Использовать небольшие значения epsilon для предотвращения делений на ноль и логарифмов от нуля.
  3. Ограничивать диапазон градиентов и значений потерь с помощью clipByValue.
  4. Настраивать learning rate, чтобы избежать резких перепадов значений.
  5. Использовать tf.where для безопасной замены некорректных значений на допустимые.

Эти методы обеспечивают устойчивость обучения и корректное обновление весов модели, предотвращая критические ошибки из-за NaN и Inf.