tf.compat.v1.gather
Извлечение фрагментов из параметра params по оси axis в соответствии с индексами. (устаревшие аргументы)
tf.compat.v1.gather(
params, indices, validate_indices=None, name=None, axis=None, batch_dims=0
)
Извлечение фрагментов из params по оси axis в соответствии с indices. indices должен быть целочисленным тензором любой размерности (часто 1-мерным).
Tensor.getitem работает для скаляров, tf.newaxis и python срезы
tf.gather расширяет индексирование для обработки тензоров индексов.
В самом простом случае это идентично скалярному индексированию:
params = tf.constant(['p0', 'p1', 'p2', 'p3', 'p4', 'p5']) params[3].numpy() b'p3' tf.gather(params, 3).numpy() b'p3'
Наиболее распространённый случай — передача тензора индексов с одной осью (это нельзя выразить как python-срез, так как индексы не последовательные):
indices = [2, 0, 2, 5] tf.gather(params, indices).numpy() array([b'p2', b'p0', b'p2', b'p5'], dtype=object)
Индексы могут иметь любую форму. Когда у params одна ось, форма результата равна форме входных данных:
tf.gather(params, [[2, 0], [2, 5]]).numpy()
array([[b'p2', b'p0'],
[b'p2', b'p5']], dtype=object)
У params также может быть любая форма. gather может выбирать фрагменты по любой оси в зависимости от аргумента axis (который по умолчанию равен 0). Ниже он используется для извлечения первых строк, а затем столбцов из матрицы:
params = tf.constant([[0, 1.0, 2.0],
[10.0, 11.0, 12.0],
[20.0, 21.0, 22.0],
[30.0, 31.0, 32.0]])
tf.gather(params, indices=[3,1]).numpy()
array([[30., 31., 32.],
[10., 11., 12.]], dtype=float32)
tf.gather(params, indices=[2,1], axis=1).numpy()
array([[ 2., 1.],
[12., 11.],
[22., 21.],
[32., 31.]], dtype=float32)
В общем случае: форма результата имеет ту же форму, что и входные данные, с осью индексирования, заменённой формой индексов.
def result_shape(p_shape, i_shape, axis=0): return p_shape[:axis] + i_shape + p_shape[axis+1:] result_shape([1, 2, 3], [], axis=1) [1, 3] result_shape([1, 2, 3], [7], axis=1) [1, 7, 3] result_shape([1, 2, 3], [7, 5], axis=1) [1, 7, 5, 3]
Вот несколько примеров:
params.shape.as_list() [4, 3] indices = tf.constant([[0, 2]]) tf.gather(params, indices=indices, axis=0).shape.as_list() [1, 2, 3] tf.gather(params, indices=indices, axis=1).shape.as_list() [4, 1, 2]
params = tf.random.normal(shape=(5, 6, 7, 8)) indices = tf.random.uniform(shape=(10, 11), maxval=7, dtype=tf.int32) result = tf.gather(params, indices, axis=2) result.shape.as_list() [5, 6, 10, 11, 8]
Потому что каждый индекс берёт срез из params, и помещает его в соответствующее место в выходных данных. Для примера выше
# For any location in indices
a, b = 0, 1
tf.reduce_all(
# the corresponding slice of the result
result[:, :, a, b, :] ==
# is equal to the slice of `params` along `axis` at the index.
params[:, :, indices[a, b], :]
).numpy()
True
Группировка:
Аргумент batch_dims позволяет собирать разные элементы из каждого элемента пакета.
Использование batch_dims=1 эквивалентно внешнему циклу по первой оси params и indices:
params = tf.constant([
[0, 0, 1, 0, 2],
[3, 0, 0, 0, 4],
[0, 5, 0, 6, 0]])
indices = tf.constant([
[2, 4],
[0, 4],
[1, 3]])
tf.gather(params, indices, axis=1, batch_dims=1).numpy()
array([[1, 2],
[3, 4],
[5, 6]], dtype=int32)
Эквивалентно:
def manually_batched_gather(params, indices, axis):
batch_dims=1
result = []
for p,i in zip(params, indices):
r = tf.gather(p, i, axis=axis-batch_dims)
result.append(r)
return tf.stack(result)
manually_batched_gather(params, indices, axis=1).numpy()
array([[1, 2],
[3, 4],
[5, 6]], dtype=int32)
Более высокие значения batch_dims эквивалентны множественным вложенным циклам по внешним осям params и indices. Поэтому функция общей формы:
def batched_result_shape(p_shape, i_shape, axis=0, batch_dims=0):
return p_shape[:axis] + i_shape[batch_dims:] + p_shape[axis+1:]
batched_result_shape(
p_shape=params.shape.as_list(),
i_shape=indices.shape.as_list(),
axis=1,
batch_dims=1)
[3, 2]
tf.gather(params, indices, axis=1, batch_dims=1).shape.as_list() [3, 2]
Это естественно возникает, если вам нужно использовать индексы операции, такой как tf.argsort или tf.math.top_k, где последний размер индексов индексирует последний размер входа в соответствующем месте. В этом случае вы можете использовать tf.gather(values, indices, batch_dims=-1).
См. также:
-
tf.Tensor.getitem: Прямая операция индексации тензора (t[]), обрабатывает скаляры и python-срезыtensor[..., 7, 1:-1] -
tf.scatter: Набор операций, аналогичных__setitem__(t[i] = x) -
tf.gather_nd: Операция, похожая наtf.gather, но собирает сразу по нескольким осям (она может собирать элементы матрицы вместо строк или столбцов) -
tf.boolean_mask,tf.where: Бинарная индексация. -
tf.sliceиtf.strided_slice: Для доступа к реализации python-срез обработки__getitem__на более низком уровне (t[1:-1:2])
| Args | |
|---|---|
params | Тензор, из которого извлекаются значения. Должен иметь не менее axis + 1 размерности. |
indices | Индекс Tensor. Должен быть одного из следующих типов: int32, int64. Значения должны быть в диапазоне [0, params.shape[axis]). |
validate_indices | Устарело, не имеет эффекта. Индексы всегда проверяются на CPU, но никогда на GPU. |
axis | Ось Tensor. Должен быть одного из следующих типов: int32, int64. Ось в params для сбора indices из. Должен быть больше или равен batch_dims. По умолчанию — первая ось, не являющаяся осью пакета. Поддерживаются отрицательные индексы. |
batch_dims | Число осей пакета. Должен быть меньше или равен rank(indices). |
name | Имя операции (необязательно). |
| Возвращает | |
|---|---|
Тензор Tensor. Имеет тот же тип, что и params. |
© 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/compat/v1/gather