tf.contrib.nn.rank_sampled_softmax_loss
Вычисляет потерю softmax с использованием адаптивного ресемплирования на основе ранга.
tf.contrib.nn.rank_sampled_softmax_loss(
weights, biases, labels, inputs, num_sampled, num_resampled, num_classes,
num_true, sampled_values, resampling_temperature, remove_accidental_hits,
partition_strategy, name=None
)
Показано, что это улучшает потерю ранжирования после обучения по сравнению с tf.nn.sampled_softmax_loss. Для описания алгоритма и некоторых экспериментальных результатов см.: TAPAS: Двухэтапное приближенное адаптивное ресемплирование для softmax.
Ресемплирование происходит в две фазы:
- На первой фазе
num_sampledклассов выбираются с помощьюtf.nn.learned_unigram_candidate_samplerили предоставленыsampled_values. Логиты рассчитываются для этих выбранных классов. Эта фаза аналогичнаtf.nn.sampled_softmax_loss. - На второй фазе
num_resampledклассов с наибольшей предсказанной вероятностью сохраняются. ВероятностиLogSumExp(logits / resampling_temperature), где сумма берется поinputs.
Параметр resampling_temperature управляет «адаптивностью» ресемплирования. При более низких температурах ресемплирование более адаптивное, потому что оно выбирает больше кандидатов, близких к предсказанным классам. Общей стратегией является уменьшение температуры по мере обучения.
См. tf.nn.sampled_softmax_loss для получения дополнительной документации по ресемплированию и типовых значений параметров по умолчанию.
Эта операция предназначена только для обучения. Как правило, это недооценка полной потери softmax.
Распространенный случай использования — использование этого метода для обучения, а вычисление полной потери softmax для оценки или вывода. В этом случае вы должны установить partition_strategy="div" для согласования двух потерь, как показано в следующем примере:
if mode == "train":
loss = rank_sampled_softmax_loss(
weights=weights,
biases=biases,
labels=labels,
inputs=inputs,
...,
partition_strategy="div")
elif mode == "eval":
logits = tf.matmul(inputs, tf.transpose(weights))
logits = tf.nn.bias_add(logits, biases)
labels_one_hot = tf.one_hot(labels, n_classes)
loss = tf.nn.softmax_cross_entropy_with_logits(
labels=labels_one_hot,
logits=logits)
| Аргументы | |
|---|---|
weights | A Tensor или PartitionedVariable формы [num_classes, dim], или список Tensor объектов, чье объединение по размеру 0 имеет форму [num_classes, dim]. Векторы вложений классов (возможно, разделенные). |
biases | A Tensor или PartitionedVariable формы [num_classes]. Смещения классов (возможно, разделенные). |
labels | A Tensor типа int64 и формы [batch_size, num_true]. Целевые классы. Обратите внимание, что этот формат отличается от аргумента labels функции nn.softmax_cross_entropy_with_logits. |
inputs | A Tensor формы [batch_size, dim]. Активации на выходе входной сети. |
num_sampled | A int. Количество классов, которые случайным образом выбираются на одну партию. |
num_resampled | A int. Количество классов, которые выбираются из num_sampled классов с использованием адаптивного алгоритма ресемплирования. Должно быть меньше num_sampled. |
num_classes | A int. Количество возможных классов. |
num_true | A int. Количество целевых классов на пример обучения. |
sampled_values | Кортеж (sampled_candidates, true_expected_count, sampled_expected_count), возвращаемый функцией *_candidate_sampler. Если None, используйте по умолчанию nn.learned_unigram_candidate_sampler. |
resampling_temperature | Скалярный Tensor с параметром температуры для адаптивного алгоритма ресемплирования. |
remove_accidental_hits | A bool. Удалять ли «случайные совпадения», когда выбранный класс равен одному из целевых классов. |
partition_strategy | Строка, определяющая стратегию разбиения, релевантную, если len(weights) > 1. В настоящее время поддерживаются "div" и "mod". См. tf.nn.embedding_lookup для получения дополнительной информации. |
name | Имя операции (необязательно). |
| Возвращает | |
|---|---|
A batch_size 1-мерный тензор потерь выборки softmax на пример. |
| Возбуждает | |
|---|---|
ValueError | Если num_sampled <= num_resampled. |
© 2020 The TensorFlow Authors. All rights reserved.
Licensed under the Creative Commons Attribution License 3.0.
Code samples licensed under the Apache 2.0 License.
https://www.tensorflow.org/versions/r1.15/api_docs/python/tf/contrib/nn/rank_sampled_softmax_loss