Spec-Zone.ru › TensorFlow 2.4

tf.where

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

Возвращает элементы, где condition является True (мультиплексирование x и y).

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

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

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

tf.compat.v1.where_v2

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

Этот оператор имеет два режима: в одном режиме предоставлены как x, так и y, в другом режиме ни то, ни другое не предоставлено. condition всегда ожидается как tf.Tensor типа bool.

Получение индексов элементов True

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

tf.where вернёт индексы condition, которые являются True, в виде двумерного тензора с формой (n, d). (Где n — количество совпадающих индексов в condition, а d — количество измерений в condition).

Индексы выводятся в порядке следования строк.

tf.where([True, False, False, True])
<tf.Tensor: shape=(2, 1), dtype=int64, numpy=
array([[0],
       [3]])>
tf.where([[True, False], [False, True]])
<tf.Tensor: shape=(2, 2), dtype=int64, numpy=
array([[0, 0],
       [1, 1]])>
tf.where([[[True, False], [False, True], [True, True]]])
<tf.Tensor: shape=(4, 3), dtype=int64, numpy=
array([[0, 0, 0],
       [0, 1, 1],
       [0, 2, 0],
       [0, 2, 1]])>

Мультиплексирование между x и y

Если x и y предоставлены (оба имеют значения, отличные от None):

tf.where выберет форму вывода из форм condition, x, и y, для которых все три формы являются совместимыми для трансляции.

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

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

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

Вместо этого

y = tf.constant(-1, dtype=tf.float32)
tf.where(y > 0, tf.sqrt(y), y)
<tf.Tensor: shape=(), dtype=float32, numpy=-1.0>

Используйте это

tf.where(y > 0, tf.sqrt(tf.where(y > 0, y, 1)), y)
<tf.Tensor: shape=(), dtype=float32, numpy=-1.0>
Аргументы
condition tf.Tensor типа bool
x Если предоставлено, тензор того же типа, что и y, и имеет форму, совместимую для трансляции с condition и y.
y Если предоставлено, тензор того же типа, что и x, и имеет форму, совместимую для трансляции с condition и x.
name Имя операции (необязательно).
Возвращаемые значения
Если x и y предоставлены: тензор Tensor того же типа, что и x и y, с формой, полученной путём трансляции из condition, x, и y. В противном случае, тензор Tensor с формой (num_true, dim_size(condition)).
Исключения
ValueError Когда ровно один из x или y не равен None, или формы не все совместимы для трансляции.

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

Spec-Zone.ru

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