Функция nn.sparse_softmax_cross_entropy_with_logits
Функция nn.sparse_softmax_cross_entropy_with_logits применяется
к необработанным выходам модели (логитам) и целочисленным меткам классов.
Первым параметром передаются labels - тензор с номерами истинных
классов, вторым - logits - тензор с логитами. Функция сама
применяет softmax к логитам и вычисляет кросс-энтропию, возвращая
потерю для каждого примера в батче.
В отличие от nn.softmax_cross_entropy_with_logits, метки здесь
задаются не в one-hot формате, а целыми числами от 0 до
num_classes - 1. Это удобно при большом числе классов, так
как не требует создавать разреженные векторы.
Синтаксис
tf.nn.sparse_softmax_cross_entropy_with_logits(
labels, logits, name=None
)
Пример
Давайте вычислим кросс-энтропию для двух примеров с тремя классами:
import tensorflow as tf
logits = tf.constant([[2.0, 1.0, 0.1], [0.5, 2.5, 0.3]])
labels = tf.constant([0, 1])
res = tf.nn.sparse_softmax_cross_entropy_with_logits(
labels=labels, logits=logits
)
print(res)
Результат выполнения кода:
tf.Tensor([0.41703 0.48989], shape=(2,), dtype=float32)
Пример
Давайте усредним полученные потери, чтобы получить одно число для всего батча:
import tensorflow as tf
logits = tf.constant([[2.0, 1.0, 0.1], [0.5, 2.5, 0.3]])
labels = tf.constant([0, 1])
losses = tf.nn.sparse_softmax_cross_entropy_with_logits(
labels=labels, logits=logits
)
res = tf.reduce_mean(losses)
print(res)
Результат выполнения кода:
tf.Tensor(0.45346, shape=(), dtype=float32)
Пример
Давайте проверим, что при уверенном правильном предсказании потеря будет близка к нулю:
import tensorflow as tf
logits = tf.constant([[10.0, 0.0, 0.0]])
labels = tf.constant([0])
res = tf.nn.sparse_softmax_cross_entropy_with_logits(
labels=labels, logits=logits
)
print(res)
Результат выполнения кода:
tf.Tensor([4.539993e-05], shape=(1,), dtype=float32)
Смотрите также
-
функцию
softmax_cross_entropy_with_logits,
которая вычисляет кросс-энтропию с one-hot метками -
функцию
softmax,
которая преобразует логиты в вероятности -
функцию
log_softmax,
которая вычисляет логарифм softmax -
функцию
sigmoid_cross_entropy_with_logits,
которая вычисляет кросс-энтропию с сигмоидой