tf.nn.compute_accidental_hits
| Просмотреть исходный код на GitHub |
Вычислить идентификаторы позиций в sampled_candidates , соответствующие true_classes.
tf.nn.compute_accidental_hits(
true_classes, sampled_candidates, num_true, seed=None, name=None
)
В Candidate Sampling данная операция позволяет практически удалить выбранные классы, которые случайно совпадают с целевыми классами. Это делается в Sampled Softmax и Sampled Logistic.
См. нашу Справочник по алгоритмам Candidate Sampling.
Мы предполагаем, что sampled_candidates являются уникальными.
Мы называем это «случайным попаданием», когда один из целевых классов совпадает с одним из выбранных классов. Эта операция сообщает о случайных попаданиях в виде троек (index, id, weight), где index представляет номер строки в true_classes, id представляет позицию в sampled_candidates, а вес — -FLOAT_MAX.
Результат этой операции должен быть передан через операцию sparse_to_dense , а затем добавлен к логарифмам выбранных классов. Это устраняет противоречивый эффект случайного выбора истинных целевых классов в качестве шумовых классов для того же примера.
| Аргументы | |
|---|---|
true_classes | A Tensor типа int64 и формы [batch_size, num_true]. Целевые классы. |
sampled_candidates | A tensor типа int64 и формы [num_sampled]. Выход sampled_candidates CandidateSampler. |
num_true | A int. Количество целевых классов на пример обучения. |
seed | A int. Семена, специфичные для операции. По умолчанию 0. |
name | Имя операции (необязательно). |
| Возвращаемые значения | |
|---|---|
indices | A Tensor типа int32 и формы [num_accidental_hits]. Значения указывают строки в true_classes. |
ids | A Tensor типа int64 и формы [num_accidental_hits]. Значения указывают позиции в sampled_candidates. |
weights | A Tensor типа float и формы [num_accidental_hits]. Каждое значение — -FLOAT_MAX. |
© 2022 The TensorFlow Authors. All rights reserved.
Licensed under the Creative Commons Attribution License 4.0.
Code samples licensed under the Apache 2.0 License.
https://www.tensorflow.org/versions/r2.9/api_docs/python/tf/nn/compute_accidental_hits