Returns true if the operator used to compute the given node contains a bool tensor
(node: torch.fx.Node)
| 466 | |
| 467 | |
| 468 | def op_contains_bool_tensor(node: torch.fx.Node) -> bool: |
| 469 | """ |
| 470 | Returns true if the operator used to compute the given node contains a bool tensor |
| 471 | """ |
| 472 | if is_tensor_node(node) and tensor_node_is_bool(node): |
| 473 | return True |
| 474 | |
| 475 | for arg_node in node.args: |
| 476 | # pyre-ignore[6] |
| 477 | if is_tensor_node(arg_node) and tensor_node_is_bool(arg_node): |
| 478 | return True |
| 479 | |
| 480 | return False |
| 481 | |
| 482 | |
| 483 | def op_contains_high_dim_tensor(node: torch.fx.Node) -> bool: |
nothing calls this directly
no test coverage detected