Функция where
Функция where выполняет условный выбор элементов.
Первым параметром передается условие в виде тензора
булевых значений. Если передан только один параметр,
функция возвращает индексы элементов, равных True.
Если переданы второй и третий параметры, функция
возвращает элементы из второго тензора там, где условие
истинно, и из третьего тензора там, где условие ложно.
Синтаксис
tf.where(condition, [x, y])
Пример
Давайте получим индексы ненулевых элементов тензора:
import tensorflow as tf
t = tf.constant([0, 1, 0, 2, 3, 0])
res = tf.where(t > 0)
print(res)
Результат выполнения кода:
tf.Tensor(
[[1]
[3]
[4]], shape=(3, 1), dtype=int64)
Пример
Давайте выберем элементы из двух тензоров по условию:
<+python+>
import tensorflow as tf
condition = tf.constant([True, False, True, False])
x = tf.constant([1, 2, 3, 4])
y = tf.constant([10, 20, 30, 40])
res = tf.where(condition, x, y)
print(res)
<-python+>
Результат выполнения кода:
tf.Tensor([ 1 20 3 40], shape=(4,), dtype=int32)
Пример
Давайте применим условие к двумерному тензору:
import tensorflow as tf
t = tf.constant([[1, 2, 3], [4, 5, 6]])
res = tf.where(t > 3, t, tf.zeros_like(t))
print(res)
Результат выполнения кода:
tf.Tensor(
[[0 0 0]
[4 5 6]], shape=(2, 3), dtype=int32)
Смотрите также
-
функцию
boolean_mask,
которая выбирает элементы по булевой маске -
функцию
gather,
которая собирает элементы по индексам -
функцию
constant,
которая создает тензор из переданных данных -
функцию
zeros,
которая создает тензор, заполненный нулями