Handle aten.conv2d: (input, weight, bias, stride, padding, dilation, groups).
(P: MLXProgramBuilder, n: Node)
| 2428 | |
| 2429 | @REGISTRY.register(target=[torch.ops.aten.conv2d.default]) |
| 2430 | def _conv2d_handler(P: MLXProgramBuilder, n: Node) -> Slot: |
| 2431 | """Handle aten.conv2d: (input, weight, bias, stride, padding, dilation, groups).""" |
| 2432 | require_args(n.args, 2, 7, "aten.conv2d") |
| 2433 | require_kwargs(P.kwargs(n), set(), "aten.conv2d") |
| 2434 | x_node, w_node = n.args[0:2] |
| 2435 | bias_node = n.args[2] if len(n.args) > 2 else None |
| 2436 | groups = n.args[6] if len(n.args) > 6 else 1 |
| 2437 | stride = _normalize_conv_param(n.args[3] if len(n.args) > 3 else [1, 1], 2, 1) |
| 2438 | padding = _normalize_conv_param(n.args[4] if len(n.args) > 4 else [0, 0], 2, 0) |
| 2439 | dilation = _normalize_conv_param(n.args[5] if len(n.args) > 5 else [1, 1], 2, 1) |
| 2440 | return _emit_conv( |
| 2441 | P, n, x_node, w_node, bias_node, stride, padding, dilation, groups, ndim=2 |
| 2442 | ) |
| 2443 | |
| 2444 | |
| 2445 | @REGISTRY.register(target=[torch.ops.aten.conv3d.default]) |
nothing calls this directly
no test coverage detected