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

Function _conv2d_handler

backends/mlx/ops.py:2430–2442  ·  view source on GitHub ↗

Handle aten.conv2d: (input, weight, bias, stride, padding, dilation, groups).

(P: MLXProgramBuilder, n: Node)

Source from the content-addressed store, hash-verified

2428
2429@REGISTRY.register(target=[torch.ops.aten.conv2d.default])
2430def _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])

Callers

nothing calls this directly

Calls 5

require_argsFunction · 0.85
require_kwargsFunction · 0.85
_normalize_conv_paramFunction · 0.85
_emit_convFunction · 0.85
kwargsMethod · 0.80

Tested by

no test coverage detected