Skip to content

Commit ad037eb

Browse files
committed
Move logical_and into MLX binary op table
1 parent 9f33375 commit ad037eb

1 file changed

Lines changed: 1 addition & 17 deletions

File tree

backends/mlx/ops.py

Lines changed: 1 addition & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -480,6 +480,7 @@ def _isnan_handler(P: MLXProgramBuilder, n: Node) -> Slot:
480480
([torch.ops.aten.minimum.default], MinimumNode, "aten.minimum", False),
481481
([torch.ops.aten.atan2.default], Atan2Node, "aten.atan2", False),
482482
([torch.ops.aten.logaddexp.default], LogAddExpNode, "aten.logaddexp", False),
483+
([torch.ops.aten.logical_and.default], LogicalAndNode, "aten.logical_and", False),
483484
([torch.ops.aten.logical_or.default], LogicalOrNode, "aten.logical_or", False),
484485
(
485486
[torch.ops.aten.bitwise_and.Tensor, torch.ops.aten.bitwise_and.Scalar],
@@ -3096,23 +3097,6 @@ def _bitwise_not_handler(P: MLXProgramBuilder, n: Node) -> Slot:
30963097
)
30973098

30983099

3099-
@REGISTRY.register(target=[torch.ops.aten.logical_and.default])
3100-
def _logical_and_handler(P: MLXProgramBuilder, n: Node) -> Slot:
3101-
"""Handle aten.logical_and on bool tensors."""
3102-
args = P.args(n)
3103-
require_args(args, 2, 2, "aten.logical_and")
3104-
require_kwargs(P.kwargs(n), set(), "aten.logical_and")
3105-
out = P.make_or_get_slot(n)
3106-
P.emit(
3107-
LogicalAndNode(
3108-
a=P.slot_to_tid(args[0]),
3109-
b=P.slot_to_tid(args[1]),
3110-
out=P.slot_to_tid(out),
3111-
)
3112-
)
3113-
return out
3114-
3115-
31163100
@REGISTRY.register(target=[torch.ops.aten.scalar_tensor.default])
31173101
def _scalar_tensor_handler(P: MLXProgramBuilder, n: Node) -> Slot:
31183102
"""This is equivalent to torch.full([], scalar, dtype=dtype)."""

0 commit comments

Comments
 (0)