Функция gather_nd
Функция gather_nd извлекает элементы или срезы из
входного тензора по многомерным индексам. Первым
параметром функция принимает исходный тензор, вторым -
тензор индексов. В отличие от gather, который
работает с индексами по одной оси, gather_nd
позволяет задавать координаты сразу по нескольким
измерениям. Третьим необязательным параметром можно
передать batch_dims для пакетной обработки.
Синтаксис
tf.gather_nd(params, indices, [batch_dims])
Пример
Давайте извлечем элемент из одномерного тензора
по индексу 2:
import tensorflow as tf
t = tf.constant([1, 2, 3, 4, 5])
res = tf.gather_nd(t, [2])
print(res)
Результат выполнения кода:
tf.Tensor(3, shape=(), dtype=int32)
Пример
Давайте извлечем несколько элементов из одномерного
тензора по индексам 0, 2 и 4:
import tensorflow as tf
t = tf.constant([1, 2, 3, 4, 5])
res = tf.gather_nd(t, [[0], [2], [4]])
print(res)
Результат выполнения кода:
tf.Tensor([1 3 5], shape=(3,), dtype=int32)
Пример
Давайте извлечем элемент из двумерного тензора
по координатам строка 1, столбец 2:
import tensorflow as tf
t = tf.constant([[1, 2, 3], [4, 5, 6]])
res = tf.gather_nd(t, [1, 2])
print(res)
Результат выполнения кода:
tf.Tensor(6, shape=(), dtype=int32)
Пример
Давайте извлечем несколько элементов из двумерного
тензора по координатам [0, 1] и [1, 2]:
import tensorflow as tf
t = tf.constant([[1, 2, 3], [4, 5, 6]])
res = tf.gather_nd(t, [[0, 1], [1, 2]])
print(res)
Результат выполнения кода:
tf.Tensor([2 6], shape=(2,), dtype=int32)
Пример
Давайте извлечем целую строку из двумерного тензора, указав индекс только по первому измерению:
import tensorflow as tf
t = tf.constant([[1, 2, 3], [4, 5, 6]])
res = tf.gather_nd(t, [[1]])
print(res)
Результат выполнения кода:
tf.Tensor([[4 5 6]], shape=(1, 3), dtype=int32)
Пример
Давайте извлечем срез из трехмерного тензора
по координатам [0, 1] и [1, 0]:
import tensorflow as tf
t = tf.constant([
[[1, 2], [3, 4]],
[[5, 6], [7, 8]]
])
res = tf.gather_nd(t, [[0, 1], [1, 0]])
print(res)
Результат выполнения кода:
tf.Tensor([[3 4]
[5 6]], shape=(2, 2), dtype=int32)
Смотрите также
-
функцию
gather,
которая извлекает срезы по индексам вдоль одной оси -
функцию
scatter_nd,
которая распределяет значения по индексам -
функцию
slice,
которая извлекает непрерывный срез тензора -
функцию
boolean_mask,
которая извлекает элементы по булевой маске