Spec-Zone.ru › TensorFlow 2.9

tf.where

Просмотреть исходный код на GitHub

Возвращает индексы ненулевых элементов или объединяет x и y.

Просмотр псевдонимов

Псевдонимы для миграции

См. Руководство по миграции для получения дополнительных сведений.

tf.compat.v1.where_v2

tf.where(
    condition, x=None, y=None, name=None
)

Данная операция имеет два режима:

  1. Возвращение индексов ненулевых элементов - Когда предоставлен только condition, результат — тензор int64, где каждая строка — индекс ненулевого элемента condition. Форма результата — [tf.math.count_nonzero(condition), tf.rank(condition)].
  2. Объединение x и y - Когда предоставлены x и y, результат имеет форму x, y, и condition вместе с применением трансляции. Результат берется из x там, где condition ненулевой, или из y там, где condition равен нулю.

1. Возвращение индексов ненулевых элементов

Примечание: В этом режиме condition может иметь тип bool или любой числовой тип.

Если x и y не предоставлены (оба равны None):

tf.where вернёт индексы condition , которые не равны нулю, в виде двумерного тензора с формой [n, d], где n — количество ненулевых элементов в condition (tf.count_nonzero(condition)), а d — количество осей condition (tf.rank(condition)).

Индексы выводятся в строчно-главном порядке. condition может иметь тип tf.bool или любой числовой dtype.

Здесь condition — тензор с 1 осью bool и 2 True значениями. Результат имеет форму [2,1]

tf.where([True, False, False, True]).numpy()
array([[0],
       [3]])

Здесь condition — двухосный целочисленный тензор с 3 ненулевыми значениями. Результат имеет форму [3, 2].

tf.where([[1, 0, 0], [1, 0, 1]]).numpy()
array([[0, 0],
       [1, 0],
       [1, 2]])

Здесь condition — трёххосный плавающий тензор с 5 ненулевыми значениями. Форма результата [5, 3].

float_tensor = [[[0.1, 0], [0, 2.2], [3.5, 1e6]],
                [[0,   0], [0,   0], [99,    0]]]
tf.where(float_tensor).numpy()
array([[0, 0, 0],
       [0, 1, 1],
       [0, 2, 0],
       [0, 2, 1],
       [1, 2, 0]])

Эти индексы такие же, как у tf.sparse.SparseTensor при представлении тензора условия:

sparse = tf.sparse.from_dense(float_tensor)
sparse.indices.numpy()
array([[0, 0, 0],
       [0, 1, 1],
       [0, 2, 0],
       [0, 2, 1],
       [1, 2, 0]])

Комплексное число считается ненулевым, если ненулевым является либо действительная, либо мнимая часть:

tf.where([complex(0.), complex(1.), 0+1j, 1+1j]).numpy()
array([[1],
       [2],
       [3]])

2. Объединение x и y

Примечание: В этом режиме condition должен иметь тип bool.

Если x и y также предоставлены (оба имеют значения, отличные от None), то тензор condition действует как маска, выбирающая, должен ли соответствующий элемент/строка в выводе браться из x (если элемент в condition ненулевой) или из y (если он нулевой).

Форма результата формируется путём распространения форм condition, x, и y.

Когда все три входных данных имеют одинаковый размер, каждый из них обрабатывается поэлементно.

tf.where([True, False, False, True],
         [1, 2, 3, 4],
         [100, 200, 300, 400]).numpy()
array([  1, 200, 300,   4], dtype=int32)

Существует два основных правила для трансляции:

  1. Если у тензора меньше осей, чем у других, к левой части формы добавляются оси длины 1.
  2. Оси с длиной 1 растягиваются для соответствия соответствующим осям других тензоров.

Вектор длиной 1 растягивается для соответствия другим векторам:

tf.where([True, False, False, True], [1, 2, 3, 4], [100]).numpy()
array([  1, 100, 100,   4], dtype=int32)

Скаляр расширяется для соответствия другим аргументам:

tf.where([[True, False], [False, True]], [[1, 2], [3, 4]], 100).numpy()
array([[  1, 100], [100,   4]], dtype=int32)
tf.where([[True, False], [False, True]], 1, 100).numpy()
array([[  1, 100], [100,   1]], dtype=int32)

Скаляр condition возвращает полный тензор x или y, с применённой трансляцией.

tf.where(True, [1, 2, 3, 4], 100).numpy()
array([1, 2, 3, 4], dtype=int32)
tf.where(False, [1, 2, 3, 4], 100).numpy()
array([100, 100, 100, 100], dtype=int32)

Для примера трансляции, нетривиального случая, тензор condition имеет форму [3], тензор x имеет форму [3,3], а тензор y имеет форму [3,1]. Сначала трансляция расширяет форму condition до [1,3]. Конечная форма трансляции — [3,3]. condition выберет столбцы из x и y. Поскольку y имеет только один столбец, все столбцы из y будут идентичными.

tf.where([True, False, True],
         x=[[1, 2, 3],
            [4, 5, 6],
            [7, 8, 9]],
         y=[[100],
            [200],
            [300]]
).numpy()
array([[ 1, 100, 3],
       [ 4, 200, 6],
       [ 7, 300, 9]], dtype=int32)

Обратите внимание, что если градиент любого из ветвей tf.where генерирует NaN, то градиент всего tf.where будет NaN. Это происходит потому, что вычисление градиента для tf.where объединяет обе ветви по соображениям производительности.

Решение — использовать вложенный tf.where для гарантирования отсутствия асимптоты функции и предотвращения вычисления значения с градиентом NaN путём замены опасных входных значений на безопасные.

Вместо этого

x = tf.constant(0., dtype=tf.float32)
with tf.GradientTape() as tape:
  tape.watch(x)
  y = tf.where(x < 1., 0., 1. / x)
print(tape.gradient(y, x))
tf.Tensor(nan, shape=(), dtype=float32)

Хотя значения 1. / x никогда не используются, их градиент будет NaN когда x = 0. Вместо этого мы должны защитить это с помощью другого tf.where

x = tf.constant(0., dtype=tf.float32)
with tf.GradientTape() as tape:
  tape.watch(x)
  safe_x = tf.where(tf.equal(x, 0.), 1., x)
  y = tf.where(x < 1., 0., 1. / safe_x)
print(tape.gradient(y, x))
tf.Tensor(0.0, shape=(), dtype=float32)

См. также:

  • tf.sparse - Индексы, возвращаемые первой формой tf.where, могут быть полезны в объектах tf.sparse.SparseTensor.
  • tf.gather_nd, tf.scatter_nd и родственные операции - Учитывая список индексов, возвращаемых из tf.where, семейство операций scatter и gather может быть использовано для извлечения значений или вставки значений по этим индексам.
  • tf.strings.length - tf.string не является допустимым типом для condition. Используйте длину строки вместо этого.
Аргументы
condition A tf.Tensor типа bool или любого числового типа. condition должен иметь тип bool когда x и y предоставлены.
x Если предоставлен, тензор того же типа, что и y, и имеющий форму, совместимую с трансляцией с condition и y.
y Если предоставлен, тензор того же типа, что и x, и имеющий форму, совместимую с трансляцией с condition и x.
name Имя операции (необязательно).
Возвращаемое значение
Если x и y предоставлены: тензор Tensor того же типа, что и x и y, и формой, полученной трансляцией из condition, x, и y. В противном случае, тензор Tensor с формой [tf.math.count_nonzero(condition), tf.rank(condition)].
Исключения
ValueError Если ровно один из x или y не равен None, или формы несовместимы с трансляцией.

© 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/where

Spec-Zone.ru

Настройки Оффлайн Что нового Помощь О нас
Spec-Zone .ru
спецификации, руководства, описания, API