Returns true if a given node contains a tensor with bool dtype
(node: torch.fx.Node)
| 426 | |
| 427 | |
| 428 | def tensor_node_is_bool(node: torch.fx.Node) -> bool: |
| 429 | """ |
| 430 | Returns true if a given node contains a tensor with bool dtype |
| 431 | """ |
| 432 | if isinstance(node.meta["val"], FakeTensor): |
| 433 | return node.meta["val"].dtype == torch.bool |
| 434 | if isinstance(node.meta["val"], list) or isinstance(node.meta["val"], tuple): |
| 435 | for fake_tensor in node.meta["val"]: |
| 436 | if isinstance(fake_tensor, FakeTensor): |
| 437 | if fake_tensor.dtype == torch.bool: |
| 438 | return True |
| 439 | return False |
| 440 | |
| 441 | |
| 442 | def ndim_of(node: Any) -> Optional[int]: |
no outgoing calls
no test coverage detected