tf.where
Возвращает индексы ненулевых элементов или выполняет выбор между 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 — тензор bool с 1 осью и 2 True значениями. Результат имеет форму [2,1]
tf.where([True, False, False, True]).numpy()
array([[0],
[3]])Здесь condition — целочисленный тензор с 2 осями и 3 ненулевыми значениями. Результат имеет форму [3, 2].
tf.where([[1, 0, 0], [1, 0, 1]]).numpy()
array([[0, 0],
[1, 0],
[1, 2]])Здесь condition — вещественный тензор с 3 осями и 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 также предоставлены (оба имеют ненулевые значения), тензор 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 | Тензор типа bool или любого числового типа. condition должен быть типа bool, когда x и y предоставлены. |
x | Если предоставлен, тензор того же типа, что и y, и его форма совместима с формами condition и y. |
y | Если предоставлен, тензор того же типа, что и x, и его форма совместима с формами condition и x. |
name | Имя операции (необязательно). |
| Возвращаемое значение | |
|---|---|
Если x и y предоставлены: тензор того же типа, что и x и y, и форма, полученная от распространения condition, x и y. В противном случае, тензор с формой [tf.math.count_nonzero(condition), tf.rank(condition)]. |
| Исключения | |
|---|---|
ValueError | Если ровно один из x или y ненулевой, или формы несовместимы. |
© 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/api_docs/python/tf/where