Spec-Zone.ru › TensorFlow

tf.where

Возвращает индексы ненулевых элементов или выполняет выбор между x и y.

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

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

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

tf.compat.v1.where_v2

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

Используется в ноутбуках

Используется в руководстве Используется в учебниках
  • Типы расширений
  • Улучшенная производительность с tf.function
  • Строки Unicode
  • Перенос обучения и тонкая настройка
  • Интегрированные градиенты
  • Существенное неустановленное заражение облегчает быстрое распространение нового коронавируса (SARS-CoV2)
  • Заметки о выпуске TFP (ноутбук 0.11.0)
  • Байесовский анализ точек переключения

Эта операция имеет два режима:

  1. Возвращение индексов ненулевых элементов - Когда предоставлен только condition, результат — тензор int64, где каждая строка — индекс ненулевого элемента condition. Форма результата — [tf.math.count_nonzero(condition), tf.rank(condition)].
  2. Выбор между 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.
  2. Оси длиной 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

Spec-Zone.ru

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