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