tf.case
| Просмотреть исходный код на GitHub |
Создать операцию выбора.
tf.case(
pred_fn_pairs, default=None, exclusive=False, strict=False, name='case'
)
См. также tf.switch_case.
Параметр pred_fn_pairs представляет собой словарь или список пар размера N. Каждая пара содержит булеву скалярную тензор и python-вызов, который создаёт тензоры, возвращаемые, если булево значение равно True. default — это вызов, генерирующий список тензоров. Все вызовы в pred_fn_pairs, а также default (если указано), должны возвращать одинаковое количество и типы тензоров.
Если exclusive==True, все предикаты оцениваются, и выбрасывается исключение, если более одного из предикатов оценивается как True. Если exclusive==False, выполнение останавливается на первом предикате, который оценивается как True, и возвращаются тензоры, сгенерированные соответствующей функцией. Если ни один из предикатов не оценивается как True, эта операция возвращает тензоры, сгенерированные default.
tf.case поддерживает вложенные структуры, как реализовано в tf.contrib.framework.nest. Все вызовы должны возвращать одну и ту же (возможно, вложенную) структуру значений, списки, кортежи и/или именованные кортежи. Исключениями являются только одиночные списки и кортежи: при возвращении вызовом они неявно распаковываются в одиночные значения. Это поведение отключается путём передачи strict=True.
Если для pred_fn_pairs используется неупорядоченный словарь, порядок проверок условия не гарантируется. Однако порядок гарантированно является детерминированным, так что переменные, созданные в условных ветвях, создаются в фиксированном порядке в каждом запуске.
Пример 1:
Псевдокод:
if (x < y) return 17; else return 23;
Выражения:
f1 = lambda: tf.constant(17) f2 = lambda: tf.constant(23) r = tf.case([(tf.less(x, y), f1)], default=f2)
Пример 2:
Псевдокод:
if (x < y && x > z) raise OpError("Only one predicate may evaluate to True");
if (x < y) return 17;
else if (x > z) return 23;
else return -1;
Выражения:
def f1(): return tf.constant(17)
def f2(): return tf.constant(23)
def f3(): return tf.constant(-1)
r = tf.case({tf.less(x, y): f1, tf.greater(x, z): f2},
default=f3, exclusive=True)
| Аргументы | |
|---|---|
pred_fn_pairs | Словарь или список пар булевой скалярной тензора и вызова, возвращающего список тензоров. |
default | Необязательный вызов, возвращающий список тензоров. |
exclusive | True, если разрешено, чтобы не более одного предиката оценивалось как True. |
strict | Булево значение, которое включает/выключает «строгий» режим; см. выше. |
name | Имя для этой операции (необязательно). |
| Возвращаемые значения | |
|---|---|
Тензоры, возвращаемые первой парой, предикат которой оценивается как True, или те, что возвращаются default если ни один не выполняется. |
| Исключения | |
|---|---|
TypeError | Если pred_fn_pairs не является списком/словарём. |
TypeError | Если pred_fn_pairs является списком, но не содержит 2-кортежей. |
TypeError | Если fns[i] не является вызовом для любого i, или default не является вызовом. |
Совместимость с Eager
Неупорядоченные словари не поддерживаются в режиме eager, когда exclusive=False. Используйте вместо этого список кортежей.
© 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/r1.15/api_docs/python/tf/case