Метод scatter_nd_update класса Variable
Метод scatter_nd_update класса Variable применяется
для выборочного обновления элементов переменной по указанным
индексам. Первым параметром метод принимает тензор индексов
indices, определяющий позиции элементов для обновления.
Вторым параметром передается тензор updates с новыми
значениями, которые нужно записать в указанные позиции.
Метод изменяет переменную на месте и возвращает обновленную
переменную.
Синтаксис
v.scatter_nd_update(indices, updates)
Пример
Давайте создадим переменную из пяти элементов и обновим второй и четвертый элементы новыми значениями:
import tensorflow as tf
v = tf.Variable([1, 2, 3, 4, 5])
indices = tf.constant([1, 3])
updates = tf.constant([20, 40])
v.scatter_nd_update(indices, updates)
print(v)
Результат выполнения кода:
<tf.Variable 'Variable:0' shape=(5,) dtype=int32, numpy=array([ 1, 20, 3, 40, 5], dtype=int32)>
Пример
Давайте создадим двумерную переменную и обновим элементы сразу в нескольких позициях:
Результат выполнения кода:
<tf.Variable 'Variable:0' shape=(2, 3) dtype=int32, numpy=
array([[ 1, 20, 3],
[ 4, 5, 60]], dtype=int32)>
Пример
Давайте вернем результат обновления в отдельную переменную и выведем его значение:
import tensorflow as tf
v = tf.Variable([1, 2, 3, 4, 5])
indices = tf.constant([0, 4])
updates = tf.constant([10, 50])
res = v.scatter_nd_update(indices, updates)
print(res.numpy())
Результат выполнения кода:
[10 2 3 4 50]
Смотрите также
-
класс
Variable,
который создает изменяемый тензор -
метод
scatter_update,
который обновляет срезы переменной по индексам -
метод
assign,
который присваивает переменной новое значение -
метод
assign_add,
который прибавляет значение к переменной