(P: MLXProgramBuilder, n: Node)
| 1074 | |
| 1075 | @REGISTRY.register(target=[torch.ops.aten.embedding.default]) |
| 1076 | def _embedding_handler(P: MLXProgramBuilder, n: Node) -> Slot: |
| 1077 | args = P.args(n) |
| 1078 | require_args(args, 2, 3, "aten.embedding") |
| 1079 | # "padding_idx", "scale_grad_by_freq", "sparse" are training only args |
| 1080 | # and ignored |
| 1081 | require_kwargs( |
| 1082 | P.kwargs(n), {"padding_idx", "scale_grad_by_freq", "sparse"}, "aten.embedding" |
| 1083 | ) |
| 1084 | w, x = args[0], args[1] |
| 1085 | # padding_idx (args[2] if present) is ignored - only affects gradients |
| 1086 | out = P.make_or_get_slot(n) |
| 1087 | P.emit( |
| 1088 | TakeNode( |
| 1089 | x=P.slot_to_tid(w), |
| 1090 | index=IntOrVidOrTid.from_tid(P.slot_to_tid(x)), |
| 1091 | out=P.slot_to_tid(out), |
| 1092 | axis=0, |
| 1093 | ) |
| 1094 | ) |
| 1095 | return out |
| 1096 | |
| 1097 | |
| 1098 | @REGISTRY.register(target=[torch.ops.aten.add.Tensor, torch.ops.aten.add.Scalar]) |
nothing calls this directly
no test coverage detected