tf.where
| Просмотреть исходный код на GitHub |
Возвращает индексы ненулевых элементов или объединяет x и y.
tf.where(
condition, x=None, y=None, name=None
)
Данная операция имеет два режима:
-
Возвращение индексов ненулевых элементов - Когда предоставлен только
condition, результат — тензорint64, где каждая строка — индекс ненулевого элементаcondition. Форма результата —[tf.math.count_nonzero(condition), tf.rank(condition)]. -
Объединение
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 растягиваются для соответствия соответствующим осям других тензоров.
Вектор длиной 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