tf.gather_nd
| Просмотреть исходный код на GitHub |
Извлечение фрагментов из params в тензор с формой, заданной indices.
tf.gather_nd(
params, indices, batch_dims=0, name=None
)
indices — это K-мерный целочисленный тензор, который лучше всего рассматривать как (K-1)-мерный тензор индексов в params, где каждый элемент определяет фрагмент params.
output[\\(i_0, ..., i_{K-2}\\)] = params[indices[\\(i_0, ..., i_{K-2}\\)]]
В то время как в tf.gather indices определяет фрагменты в первом измерении params, в tf.gather_nd, indices определяет фрагменты в первых N измерениях params, где N = indices.shape[-1].
Последнее измерение indices может быть не больше ранга params.
indices.shape[-1] <= params.rank
Последнее измерение indices соответствует элементам (если indices.shape[-1] == params.rank) или фрагментам (если indices.shape[-1] < params.rank) вдоль измерения indices.shape[-1] params. Тензор результата имеет форму
indices.shape[:-1] + params.shape[indices.shape[-1]:]
Кроме того, и 'params', и 'indices' могут иметь M ведущих пакетных измерений, которые точно совпадают. В этом случае 'batch_dims' должен быть равен M.
Обратите внимание, что на процессоре CPU, если найден вне диапазона индекс, возвращается ошибка. На GPU, если найден индекс вне диапазона, в соответствующее значение выходного тензора записывается 0.
Ниже приведены некоторые примеры.
Простой индексирование матрицы:
indices = [[0, 0], [1, 1]] params = [['a', 'b'], ['c', 'd']] output = ['a', 'd']
Индексирование фрагментами матрицы:
indices = [[1], [0]] params = [['a', 'b'], ['c', 'd']] output = [['c', 'd'], ['a', 'b']]
Индексирование 3-тензора:
indices = [[1]]
params = [[['a0', 'b0'], ['c0', 'd0']],
[['a1', 'b1'], ['c1', 'd1']]]
output = [[['a1', 'b1'], ['c1', 'd1']]]
indices = [[0, 1], [1, 0]]
params = [[['a0', 'b0'], ['c0', 'd0']],
[['a1', 'b1'], ['c1', 'd1']]]
output = [['c0', 'd0'], ['a1', 'b1']]
indices = [[0, 0, 1], [1, 0, 1]]
params = [[['a0', 'b0'], ['c0', 'd0']],
[['a1', 'b1'], ['c1', 'd1']]]
output = ['b0', 'b1']
Примеры ниже предназначены для случая, когда только у индексов есть дополнительные ведущие измерения. Если у 'params' и 'indices' есть дополнительные ведущие пакетные измерения, используйте параметр 'batch_dims' для запуска gather_nd в пакетном режиме.
Пакетное индексирование матрицы:
indices = [[[0, 0]], [[0, 1]]] params = [['a', 'b'], ['c', 'd']] output = [['a'], ['b']]
Пакетное индексирование фрагментами матрицы:
indices = [[[1]], [[0]]] params = [['a', 'b'], ['c', 'd']] output = [[['c', 'd']], [['a', 'b']]]
Пакетное индексирование 3-тензора:
indices = [[[1]], [[0]]]
params = [[['a0', 'b0'], ['c0', 'd0']],
[['a1', 'b1'], ['c1', 'd1']]]
output = [[[['a1', 'b1'], ['c1', 'd1']]],
[[['a0', 'b0'], ['c0', 'd0']]]]
indices = [[[0, 1], [1, 0]], [[0, 0], [1, 1]]]
params = [[['a0', 'b0'], ['c0', 'd0']],
[['a1', 'b1'], ['c1', 'd1']]]
output = [[['c0', 'd0'], ['a1', 'b1']],
[['a0', 'b0'], ['c1', 'd1']]]
indices = [[[0, 0, 1], [1, 0, 1]], [[0, 1, 1], [1, 1, 0]]]
params = [[['a0', 'b0'], ['c0', 'd0']],
[['a1', 'b1'], ['c1', 'd1']]]
output = [['b0', 'b1'], ['d0', 'c1']]
Примеры с пакетными 'params' и 'indices':
batch_dims = 1
indices = [[1], [0]]
params = [[['a0', 'b0'], ['c0', 'd0']],
[['a1', 'b1'], ['c1', 'd1']]]
output = [['c0', 'd0'], ['a1', 'b1']]
batch_dims = 1
indices = [[[1]], [[0]]]
params = [[['a0', 'b0'], ['c0', 'd0']],
[['a1', 'b1'], ['c1', 'd1']]]
output = [[['c0', 'd0']], [['a1', 'b1']]]
batch_dims = 1
indices = [[[1, 0]], [[0, 1]]]
params = [[['a0', 'b0'], ['c0', 'd0']],
[['a1', 'b1'], ['c1', 'd1']]]
output = [['c0'], ['b1']]
См. также tf.gather.
| Аргументы | |
|---|---|
params | A Tensor. Тензор, из которого извлекаются значения. |
indices | A Tensor. Должен быть одного из следующих типов: int32, int64. Тензор индексов. |
name | Имя операции (необязательно). |
batch_dims | Целое число или скалярный 'Tensor'. Количество пакетных измерений. |
| Возвращаемое значение | |
|---|---|
A Tensor. Имеет тот же тип, что и params. |
© 2020 The TensorFlow Authors. All rights reserved.
Licensed under the Creative Commons Attribution License 3.0.
Code samples licensed under the Apache 2.0 License.
https://www.tensorflow.org/versions/r2.3/api_docs/python/tf/gather_nd