(P: MLXProgramBuilder, n: Node)
| 1270 | target=[torch.ops.aten.slice.Tensor, torch.ops.aten.slice_copy.Tensor] |
| 1271 | ) |
| 1272 | def _slice_handler(P: MLXProgramBuilder, n: Node) -> Slot: |
| 1273 | args = P.args(n) |
| 1274 | require_args(args, 4, 5, "aten.slice") |
| 1275 | require_kwargs(P.kwargs(n), set(), "aten.slice") |
| 1276 | x, dim, start, stop = args[0], args[1], args[2], args[3] |
| 1277 | step = args[4] if len(args) > 4 else 1 |
| 1278 | if start is None: |
| 1279 | start = 0 |
| 1280 | require_static_int(step, "step", "aten.slice") |
| 1281 | assert step >= 1, f"aten.slice: step must be >= 1, got {step}" |
| 1282 | out = P.make_or_get_slot(n) |
| 1283 | P.emit( |
| 1284 | SliceNode( |
| 1285 | x=P.slot_to_tid(x), |
| 1286 | out=P.slot_to_tid(out), |
| 1287 | axis=P.to_int_or_vid(dim), |
| 1288 | start=P.to_int_or_vid(start), |
| 1289 | stop=P.to_int_or_vid(stop), |
| 1290 | step=step, |
| 1291 | ) |
| 1292 | ) |
| 1293 | return out |
| 1294 | |
| 1295 | |
| 1296 | @REGISTRY.register(target=[torch.ops.aten.narrow.default]) |
nothing calls this directly
no test coverage detected