tf.tensor_scatter_nd_update
| Просмотреть исходный код на GitHub |
«Разброс updates в существующий тензор в соответствии с indices.
tf.tensor_scatter_nd_update(
tensor, indices, updates, name=None
)
Эта операция создает новый тензор путём применения разреженного updates к входному tensor. Это похоже на присвоение по индексу.
# Not implemented: tensors cannot be updated inplace. tensor[indices] = updates
Если на процессоре CPU обнаружен индекс за пределами допустимого диапазона, возвращается ошибка.
- Если обнаружен индекс за пределами допустимого диапазона, индекс игнорируется.
- Порядок применения обновлений не определен, поэтому результат будет не определен, если
indicesсодержит дубликаты.
Эта операция очень похожа на tf.scatter_nd, за исключением того, что обновления разбрасываются по существующему тензору (в отличие от нулевого тензора). Если память для существующего тензора не может быть повторно использована, создается и обновляется копия.
В общем случае:
-
indices— это целочисленный тензор — индексы для обновления вtensor. -
indicesимеет по крайней мере две оси, а последняя ось — глубина векторов индексов. - Для каждого вектора индексов в
indicesсуществует соответствующая запись вupdates. - Если длина векторов индексов соответствует рангу
tensor, то каждый вектор индексов указывает на скаляры вtensor, и каждое обновление — это скаляр. - Если длина векторов индексов меньше ранга
tensor, то каждый вектор индексов указывает на срезыtensor, и форма обновлений должна соответствовать этому срезу.
В итоге это приводит к следующим ограничениям на форму:
assert tf.rank(indices) >= 2 index_depth = indices.shape[-1] batch_shape = indices.shape[:-1] assert index_depth <= tf.rank(tensor) outer_shape = tensor.shape[:index_depth] inner_shape = tensor.shape[index_depth:] assert updates.shape == batch_shape + inner_shape
Типичное использование часто намного проще, чем эта общая форма, и его лучше понимать, начиная с простых примеров:
Обновления скалярами
Самое простое использование — вставка скалярных элементов в тензор по индексу. В этом случае index_depth должен быть равен рангу входного tensor, срез каждой колонки indices — это индекс в оси входного tensor.
В этом самом простом случае ограничения на форму таковы:
num_updates, index_depth = indices.shape.as_list() assert updates.shape == [num_updates] assert index_depth == tf.rank(tensor)`
Например, чтобы вставить 4 рассеянных элемента в тензор ранга 1 с 8 элементами.
Эта операция разброса выглядит следующим образом:
tensor = [0, 0, 0, 0, 0, 0, 0, 0] # tf.rank(tensor) == 1 indices = [[1], [3], [4], [7]] # num_updates == 4, index_depth == 1 updates = [9, 10, 11, 12] # num_updates == 4 print(tf.tensor_scatter_nd_update(tensor, indices, updates)) tf.Tensor([ 0 9 0 10 11 0 0 12], shape=(8,), dtype=int32)
Длина (первая ось) updates должна быть равна длине indices: num_updates. Это число вставляемых обновлений. Каждое скалярное обновление вставляется в tensor в указанном месте.
Для тензора входного ранга tensor скалярные обновления можно вставить, используя index_depth, которое соответствует tf.rank(tensor):
tensor = [[1, 1], [1, 1], [1, 1]] # tf.rank(tensor) == 2
indices = [[0, 1], [2, 0]] # num_updates == 2, index_depth == 2
updates = [5, 10] # num_updates == 2
print(tf.tensor_scatter_nd_update(tensor, indices, updates))
tf.Tensor(
[[ 1 5]
[ 1 1]
[10 1]], shape=(3, 2), dtype=int32)
Обновления слайсами
Когда входной tensor имеет более одной оси, можно использовать операцию разброса для обновления целых слайсов.
В этом случае полезно рассматривать входной tensor как двухуровневый массив массивов. Форма этого двухуровневого массива разделена на outer_shape и inner_shape.
indices индексирует внешний уровень входного тензора (outer_shape) и заменяет подмассив в этом месте соответствующим элементом из списка updates . Форма каждого обновления — inner_shape.
При обновлении списка слайсов ограничения на форму таковы:
num_updates, index_depth = indices.shape.as_list() inner_shape = tensor.shape[:index_depth] outer_shape = tensor.shape[index_depth:] assert updates.shape == [num_updates, inner_shape]
Например, чтобы обновить строки (6, 3) tensor:
tensor = tf.zeros([6, 3], dtype=tf.int32)
Используйте глубину индексов равную единице.
indices = tf.constant([[2], [4]]) # num_updates == 2, index_depth == 1 num_updates, index_depth = indices.shape.as_list()
outer_shape это 6, внутренняя форма — 3:
outer_shape = tensor.shape[:index_depth] inner_shape = tensor.shape[index_depth:]
Индексируются 2 строки, поэтому нужно предоставить 2 updates. Каждое обновление должно соответствовать форме inner_shape.
# num_updates == 2, inner_shape==3
updates = tf.constant([[1, 2, 3],
[4, 5, 6]])
В итоге это дает:
tf.tensor_scatter_nd_update(tensor, indices, updates).numpy()
array([[0, 0, 0],
[0, 0, 0],
[1, 2, 3],
[0, 0, 0],
[4, 5, 6],
[0, 0, 0]], dtype=int32)
Примеры дополнительных обновлений слайсами
Тензор, представляющий пакет видеоклипов одинакового размера, естественным образом имеет 5 осей: [batch_size, time, width, height, channels].
Например:
batch_size, time, width, height, channels = 13,11,7,5,3 video_batch = tf.zeros([batch_size, time, width, height, channels])
Чтобы заменить выборку видеоклипов:
- Используйте глубину индекса 1 (индексируя
outer_shape:[batch_size]) - Предоставьте обновления, каждое с формой, соответствующей форме
inner_shape:[time, width, height, channels].
Чтобы заменить первые два клипа единицами:
indices = [[0],[1]] new_clips = tf.ones([2, time, width, height, channels]) tf.tensor_scatter_nd_update(video_batch, indices, new_clips)
Чтобы заменить выборку кадров в видеоклипах:
-
indicesдолжен иметь глубину индекса 2 дляouter_shape:[batch_size, time]. -
updatesдолжна иметь форму списка изображений. Каждое обновление должно иметь форму, соответствующуюinner_shape:[width, height, channels].
Чтобы заменить первый кадр первых трёх видеоклипов:
indices = [[0, 0], [1, 0], [2, 0]] # num_updates=3, index_depth=2 new_images = tf.ones([ # num_updates=3, inner_shape=(width, height, channels) 3, width, height, channels]) tf.tensor_scatter_nd_update(video_batch, indices, new_images)
Сложенные индексы
В простых случаях удобно рассматривать indices и updates как списки, но это не строгое требование. Вместо плоского num_updates, indices и updates можно сложить в batch_shape. Этот batch_shape включает все оси indices, кроме самой внутренней index_depth оси.
index_depth = indices.shape[-1] batch_shape = indices.shape[:-1]
Примечание: Единственное исключение состоит в том, чтоbatch_shapeне может быть[]. Вы не можете обновить один индекс, передав индексы с формой[index_depth].
updates должен иметь соответствующую форму batch_shape (оси до inner_shape).
assert updates.shape == batch_shape + inner_shape
Примечание: Результат эквивалентен уплощению осейbatch_shapeтензоровindicesиupdates. Это обобщение просто позволяет избежать необходимости в переформатировании, когда удобнее строить «сложенные» индексы и обновления.
С этим обобщением полные ограничения на форму таковы:
assert tf.rank(indices) >= 2 index_depth = indices.shape[-1] batch_shape = indices.shape[:-1] assert index_depth <= tf.rank(tensor) outer_shape = tensor.shape[:index_depth] inner_shape = tensor.shape[index_depth:] assert updates.shape == batch_shape + inner_shape
Например, чтобы нарисовать X на матрице (5,5) начните с этих индексов:
tensor = tf.zeros([5,5]) indices = tf.constant([ [[0,0], [1,1], [2,2], [3,3], [4,4]], [[0,4], [1,3], [2,2], [3,1], [4,0]], ]) indices.shape.as_list() # batch_shape == [2, 5], index_depth == 2 [2, 5, 2]
Здесь indices не имеют формы [num_updates, index_depth], а формы batch_shape+[index_depth].
Поскольку index_depth равна рангу tensor:
-
outer_shapeэто(5,5) -
inner_shapeэто()— каждое обновление — скаляр -
updates.shapeэтоbatch_shape + inner_shape == (5,2) + ()
updates = [ [1,1,1,1,1], [1,1,1,1,1], ]
Соединяя всё вместе:
tf.tensor_scatter_nd_update(tensor, indices, updates).numpy()
array([[1., 0., 0., 0., 1.],
[0., 1., 0., 1., 0.],
[0., 0., 1., 0., 0.],
[0., 1., 0., 1., 0.],
[1., 0., 0., 0., 1.]], dtype=float32)
| Аргументы | |
|---|---|
tensor | Тензор для копирования/обновления. |
indices | Индексы для обновления. |
updates | Обновления, которые нужно применить по индексам. |
name | Необязательное имя операции. |
| Возвращаемое значение | |
|---|---|
| Новый тензор с заданной формой и примененными обновлениями в соответствии с индексами. |
© 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.4/api_docs/python/tf/tensor_scatter_nd_update