tf.math.approx_min_k
Возвращает минимальные k значения и их индексы входного operand приближенным способом.
tf.math.approx_min_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 | Массив для поиска min-k. Должен быть типа с плавающей точкой. |
k | Указывает количество min-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_min_k с jit. См. следующий пример для поиска ближайших соседей по расстоянию l2 в квадрате:
import tensorflow as tf
@tf.function(jit_compile=True)
def l2_ann(qy, db, half_db_norms, k=10, recall_target=0.95):
dists = half_db_norms - tf.einsum('ik,jk->ij', qy, db)
return tf.nn.approx_min_k(dists, k=k, recall_target=recall_target)
qy = tf.random.uniform((256,128))
db = tf.random.uniform((2048,128))
half_db_norms = tf.norm(db, axis=1) / 2
dists, neighbors = l2_ann(qy, db, half_db_norms)В приведенном выше примере мы вычисляем db_norms/2 - dot(qy, db^T) вместо qy^2 - 2 dot(qy, db^T) + db^2 по соображениям производительности. Первый вариант использует меньше арифметических операций и генерирует тот же набор соседей.
© 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_min_k