tf.math.approx_max_k
Возвращает приблизительно максимальные k значения и их индексы входного operand.
tf.math.approx_max_k(
operand,
k,
reduction_dimension=-1,
recall_target=0.95,
reduction_input_size_override=-1,
aggregate_to_topk=True,
name=None
)
См. https://arxiv.org/abs/2206.14286 для получения подробной информации об алгоритме. Этот оператор в настоящее время оптимизирован только на TPU.
| Аргументы | |
|---|---|
operand | Массив для поиска максимальных значений. Должен быть типа с плавающей точкой. |
k | Указывает количество максимальных значений. |
reduction_dimension | Целочисленная размерность, по которой выполняется поиск. По умолчанию: -1. |
recall_target | Цель точности приближения. |
reduction_input_size_override | При установке в положительное значение переопределяет размер, определяемый operand[reduction_dim] для оценки точности. Этот параметр полезен, когда заданный operand представляет собой только подмножество общего вычисления в SPMD или распределённых конвейерах, где истинный размер входных данных не может быть определён по форме operand. |
aggregate_to_topk | Если True, агрегирует приблизительные результаты до top-k. Если False, возвращает приблизительные результаты. Количество приблизительных результатов определяется реализацией и больше или равно заданному значению k. |
name | Необязательное имя операции. |
| Возвращаемые значения | |
|---|---|
Кортеж из двух массивов. Массивы содержат максимальные k значения и соответствующие индексы вдоль reduction_dimension входного operand. Размеры массивов такие же, как у входного operand, за исключением размерности reduction_dimension: если aggregate_to_topk равно true, размерность уменьшения равна k; в противном случае она больше или равна k, где размер определяется реализацией. |
Мы рекомендуем пользователям обернуть approx_max_k с помощью jit. Следующий пример демонстрирует поиск максимального внутреннего произведения (MIPS):
import tensorflow as tf
@tf.function(jit_compile=True)
def mips(qy, db, k=10, recall_target=0.95):
dists = tf.einsum('ik,jk->ij', qy, db)
# returns (f32[qy_size, k], i32[qy_size, k])
return tf.nn.approx_max_k(dists, k=k, recall_target=recall_target)
qy = tf.random.uniform((256,128))
db = tf.random.uniform((2048,128))
dot_products, neighbors = mips(qy, db, k=20)
© 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/api_docs/python/tf/math/approx_max_k