@@ -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 ])
31173101def _scalar_tensor_handler (P : MLXProgramBuilder , n : Node ) -> Slot :
31183102 """This is equivalent to torch.full([], scalar, dtype=dtype)."""
0 commit comments