(P: MLXProgramBuilder, n: "torch.fx.node.Node")
| 5682 | |
| 5683 | @REGISTRY.register(target=[torch.ops.mlx_test.vadd.default]) |
| 5684 | def _vadd_handler(P: MLXProgramBuilder, n: "torch.fx.node.Node") -> Slot: |
| 5685 | args = P.args(n) |
| 5686 | a_slot, b_slot = args[0], args[1] |
| 5687 | out = P.make_or_get_slot(n) |
| 5688 | |
| 5689 | a_meta = n.args[0].meta.get("val") |
| 5690 | numel = a_meta.numel() |
| 5691 | dtype_int = torch_dtype_to_scalar_type(a_meta.dtype) |
| 5692 | |
| 5693 | P.emit( |
| 5694 | MetalKernelNode( |
| 5695 | name="vadd", |
| 5696 | source=vadd_source, |
| 5697 | inputs=[P.slot_to_tid(a_slot), P.slot_to_tid(b_slot)], |
| 5698 | outputs=[P.slot_to_tid(out)], |
| 5699 | grid=[ |
| 5700 | IntOrVid.from_literal(numel), |
| 5701 | IntOrVid.from_literal(1), |
| 5702 | IntOrVid.from_literal(1), |
| 5703 | ], |
| 5704 | threadgroup=[ |
| 5705 | IntOrVid.from_literal(256), |
| 5706 | IntOrVid.from_literal(1), |
| 5707 | IntOrVid.from_literal(1), |
| 5708 | ], |
| 5709 | input_names=["a", "b"], |
| 5710 | output_names=["out"], |
| 5711 | output_shapes_flat=[IntOrVid.from_literal(d) for d in a_meta.shape], |
| 5712 | output_shape_lengths=[len(a_meta.shape)], |
| 5713 | output_dtypes=[dtype_int], |
| 5714 | ) |
| 5715 | ) |
| 5716 | return out |
| 5717 | |
| 5718 | |
| 5719 | _register_vadd_handler() |
nothing calls this directly
no test coverage detected