tf.compat.v2.gather_nd
Извлечение слайсов из params в тензор с формой, заданной indices.
tf.compat.v2.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/r1.15/api_docs/python/tf/compat/v2/gather_nd