Функция nn.nce_loss
Функция nce_loss вычисляет потерю Noise Contrastive Estimation.
Она применяется для обучения моделей с большим количеством классов,
когда вычисление полной функции softmax слишком затратно.
Функция преобразует задачу многоклассовой классификации
в задачу бинарной классификации между правильным классом
и случайно выбранными отрицательными примерами.
Первым параметром передается тензор весов weights,
вторым - тензор смещений biases,
третьим - входные данные inputs,
четвертым - метки правильных классов labels,
пятым - количество отрицательных примеров num_sampled.
Также можно передать количество классов num_classes
и другие необязательные параметры.
Синтаксис
tf.nn.nce_loss(
weights,
biases,
labels,
inputs,
num_sampled,
num_classes,
[num_true],
[sampled_values],
[remove_accidental_hits],
[name]
)
Пример
Давайте вычислим потерю NCE для небольшого набора данных. Создадим веса, смещения, входы и метки:
import tensorflow as tf
tf.random.set_seed(0)
weights = tf.random.normal([5, 4])
biases = tf.zeros([5])
inputs = tf.random.normal([3, 4])
labels = tf.constant([0, 2, 4])
loss = tf.nn.nce_loss(
weights=weights,
biases=biases,
labels=labels,
inputs=inputs,
num_sampled=2,
num_classes=5
)
print(loss)
Результат выполнения кода:
tf.Tensor([...], shape=(3,), dtype=float32)
Пример
Давайте вычислим среднее значение потери NCE:
import tensorflow as tf
tf.random.set_seed(0)
weights = tf.random.normal([5, 4])
biases = tf.zeros([5])
inputs = tf.random.normal([3, 4])
labels = tf.constant([0, 2, 4])
loss = tf.nn.nce_loss(
weights=weights,
biases=biases,
labels=labels,
inputs=inputs,
num_sampled=2,
num_classes=5
)
res = tf.reduce_mean(loss)
print(res)
Результат выполнения кода:
tf.Tensor(..., shape=(), dtype=float32)
Смотрите также
-
функцию
sampled_softmax_loss,
которая вычисляет выборочную потерю softmax -
функцию
softmax_cross_entropy_with_logits,
которая вычисляет перекрестную энтропию с softmax -
функцию
sparse_softmax_cross_entropy_with_logits,
которая вычисляет разреженную перекрестную энтропию с softmax -
функцию
sigmoid_cross_entropy_with_logits,
которая вычисляет перекрестную энтропию с сигмоидой