tf.gather_nd
| Просмотреть исходный код на GitHub |
Извлечь слайсы из params в тензор с формой, заданной indices.
tf.gather_nd(
params, indices, batch_dims=0, name=None
)
indices — это Tensor индексов в params. Векторы индексов расположены вдоль последней оси indices.
Это аналогично tf.gather, в котором indices определяет слайсы в первом измерении params. В tf.gather_nd, indices определяет слайсы в первых N измерениях params, где N = indices.shape[-1].
Сбор скаляров
В самом простом случае векторы в indices индексируют весь ранг params:
tf.gather_nd(
indices=[[0, 0],
[1, 1]],
params = [['a', 'b'],
['c', 'd']]).numpy()
array([b'a', b'd'], dtype=object)
В этом случае результат имеет на 1 ось меньше, чем indices, и каждый вектор индекса заменяется скаляром, индексированным из params.
В этом случае соотношение форм:
index_depth = indices.shape[-1] assert index_depth == params.shape.rank result_shape = indices.shape[:-1]
Если indices имеет ранг K, удобно рассматривать indices как (K-1)-мерный тензор индексов в params.
Сбор слайсов
Если векторы индексов не индексируют весь ранг params , то каждое место в результате содержит слайс из params. Этот пример собирает строки из матрицы:
tf.gather_nd(
indices = [[1],
[0]],
params = [['a', 'b', 'c'],
['d', 'e', 'f']]).numpy()
array([[b'd', b'e', b'f'],
[b'a', b'b', b'c']], dtype=object)
Здесь indices содержит [2] векторы индексов, каждый длиной 1. Векторы индексов ссылаются на строки матрицы params. Каждая строка имеет форму [3], поэтому форма результата — [2, 3].
В этом случае соотношение форм:
index_depth = indices.shape[-1] outer_shape = indices.shape[:-1] assert index_depth <= params.shape.rank inner_shape = params.shape[index_depth:] output_shape = outer_shape + inner_shape
Полезно рассматривать результаты в этом случае как тензоры-тензоров. Форма внешнего тензора определяется ведущими измерениями indices. Форма внутренних тензоров — это форма одного слайса.
Пачки
Кроме того, как params, так и indices могут иметь M ведущие размерности пачек, которые точно совпадают. В этом случае batch_dims должно быть установлено в M.
Например, чтобы собрать одну строку из каждой из пачки матриц, можно установить ведущие элементы векторов индексов как их расположение в пачке:
tf.gather_nd(
indices = [[0, 1],
[1, 0],
[2, 4],
[3, 2],
[4, 1]],
params=tf.zeros([5, 7, 3])).shape.as_list()
[5, 3]
Аргумент batch_dims позволяет опустить эти ведущие размерности расположения из индекса:
tf.gather_nd(
batch_dims=1,
indices = [[1],
[0],
[4],
[2],
[1]],
params=tf.zeros([5, 7, 3])).shape.as_list()
[5, 3]
Это эквивалентно вызову отдельного gather_nd для каждого расположения в размерностях пачки.
params=tf.zeros([5, 7, 3]) indices=tf.zeros([5, 1]) batch_dims = 1 index_depth = indices.shape[-1] batch_shape = indices.shape[:batch_dims] assert params.shape[:batch_dims] == batch_shape outer_shape = indices.shape[batch_dims:-1] assert index_depth <= params.shape.rank inner_shape = params.shape[batch_dims + index_depth:] output_shape = batch_shape + outer_shape + inner_shape output_shape.as_list() [5, 3]
Дополнительные примеры
Индексирование в 3-мерный тензор:
tf.gather_nd(
indices = [[1]],
params = [[['a0', 'b0'], ['c0', 'd0']],
[['a1', 'b1'], ['c1', 'd1']]]).numpy()
array([[[b'a1', b'b1'],
[b'c1', b'd1']]], dtype=object)
tf.gather_nd(
indices = [[0, 1], [1, 0]],
params = [[['a0', 'b0'], ['c0', 'd0']],
[['a1', 'b1'], ['c1', 'd1']]]).numpy()
array([[b'c0', b'd0'],
[b'a1', b'b1']], dtype=object)
tf.gather_nd(
indices = [[0, 0, 1], [1, 0, 1]],
params = [[['a0', 'b0'], ['c0', 'd0']],
[['a1', 'b1'], ['c1', 'd1']]]).numpy()
array([b'b0', b'b1'], dtype=object)
Примеры ниже относятся к случаю, когда только у индексов есть дополнительные ведущие размерности. Если у 'params' и 'indices' есть ведущие размерности пачки, используйте параметр 'batch_dims', чтобы выполнить gather_nd в режиме пачки.
Индексирование по партиям в матрицу:
tf.gather_nd(
indices = [[[0, 0]], [[0, 1]]],
params = [['a', 'b'], ['c', 'd']]).numpy()
array([[b'a'],
[b'b']], dtype=object)
Индексирование по партиям в матрицу (слайсы):
tf.gather_nd(
indices = [[[1]], [[0]]],
params = [['a', 'b'], ['c', 'd']]).numpy()
array([[[b'c', b'd']],
[[b'a', b'b']]], dtype=object)
Индексирование по партиям в 3-мерный тензор:
tf.gather_nd(
indices = [[[1]], [[0]]],
params = [[['a0', 'b0'], ['c0', 'd0']],
[['a1', 'b1'], ['c1', 'd1']]]).numpy()
array([[[[b'a1', b'b1'],
[b'c1', b'd1']]],
[[[b'a0', b'b0'],
[b'c0', b'd0']]]], dtype=object)
tf.gather_nd(
indices = [[[0, 1], [1, 0]], [[0, 0], [1, 1]]],
params = [[['a0', 'b0'], ['c0', 'd0']],
[['a1', 'b1'], ['c1', 'd1']]]).numpy()
array([[[b'c0', b'd0'],
[b'a1', b'b1']],
[[b'a0', b'b0'],
[b'c1', b'd1']]], dtype=object)
tf.gather_nd(
indices = [[[0, 0, 1], [1, 0, 1]], [[0, 1, 1], [1, 1, 0]]],
params = [[['a0', 'b0'], ['c0', 'd0']],
[['a1', 'b1'], ['c1', 'd1']]]).numpy()
array([[b'b0', b'b1'],
[b'd0', b'c1']], dtype=object)
Примеры с пачкой 'params' и 'indices':
tf.gather_nd(
batch_dims = 1,
indices = [[1],
[0]],
params = [[['a0', 'b0'],
['c0', 'd0']],
[['a1', 'b1'],
['c1', 'd1']]]).numpy()
array([[b'c0', b'd0'],
[b'a1', b'b1']], dtype=object)
tf.gather_nd(
batch_dims = 1,
indices = [[[1]], [[0]]],
params = [[['a0', 'b0'], ['c0', 'd0']],
[['a1', 'b1'], ['c1', 'd1']]]).numpy()
array([[[b'c0', b'd0']],
[[b'a1', b'b1']]], dtype=object)
tf.gather_nd(
batch_dims = 1,
indices = [[[1, 0]], [[0, 1]]],
params = [[['a0', 'b0'], ['c0', 'd0']],
[['a1', 'b1'], ['c1', 'd1']]]).numpy()
array([[b'c0'],
[b'b1']], dtype=object)
См. также tf.gather.
| Аргументы | |
|---|---|
params | Тензор. Тензор, из которого извлекаются значения. |
indices | Тензор. Должен быть одного из следующих типов: int32, int64. Тензор индексов. |
name | Имя операции (необязательно). |
batch_dims | Целое число или скалярный тензор. Количество размерностей пачки. |
| Возвращаемое значение | |
|---|---|
Тензор. Имеет тот же тип, что и 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/gather_nd