tf.where
| Просмотреть исходный код на GitHub |
Возвращает элементы, где condition является True (мультиплексирование x и y).
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