Handle aten.conv1d: (input, weight, bias, stride, padding, dilation, groups).
(P: MLXProgramBuilder, n: Node)
| 2412 | |
| 2413 | @REGISTRY.register(target=[torch.ops.aten.conv1d.default]) |
| 2414 | def _conv1d_handler(P: MLXProgramBuilder, n: Node) -> Slot: |
| 2415 | """Handle aten.conv1d: (input, weight, bias, stride, padding, dilation, groups).""" |
| 2416 | require_args(n.args, 2, 7, "aten.conv1d") |
| 2417 | require_kwargs(P.kwargs(n), set(), "aten.conv1d") |
| 2418 | x_node, w_node = n.args[0:2] |
| 2419 | bias_node = n.args[2] if len(n.args) > 2 else None |
| 2420 | groups = n.args[6] if len(n.args) > 6 else 1 |
| 2421 | stride = _normalize_conv_param(n.args[3] if len(n.args) > 3 else 1, 1, 1) |
| 2422 | padding = _normalize_conv_param(n.args[4] if len(n.args) > 4 else 0, 1, 0) |
| 2423 | dilation = _normalize_conv_param(n.args[5] if len(n.args) > 5 else 1, 1, 1) |
| 2424 | return _emit_conv( |
| 2425 | P, n, x_node, w_node, bias_node, stride, padding, dilation, groups, ndim=1 |
| 2426 | ) |
| 2427 | |
| 2428 | |
| 2429 | @REGISTRY.register(target=[torch.ops.aten.conv2d.default]) |
nothing calls this directly
no test coverage detected