MCPcopy Create free account
hub / github.com/pytorch/executorch / _vadd_handler

Function _vadd_handler

backends/mlx/test/test_ops.py:5684–5716  ·  view source on GitHub ↗
(P: MLXProgramBuilder, n: "torch.fx.node.Node")

Source from the content-addressed store, hash-verified

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()

Callers

nothing calls this directly

Calls 7

argsMethod · 0.80
numelMethod · 0.80
emitMethod · 0.80
slot_to_tidMethod · 0.80
make_or_get_slotMethod · 0.45
getMethod · 0.45

Tested by

no test coverage detected