Skip to content

Commit ed7a23e

Browse files
authored
Channels last: preserve metadata across op replacement (pytorch#21961)
Preserve complete operator metadata when replacing ATen operators with channels-last dialect counterparts. ExportPass retracing recomputes layout-dependent shape fields while source attribution, quantization annotations, and custom backend metadata survive. Document that callers must reject or remap semantic metadata tied to dimensions changed by the replacement specification, such as per-channel quantization axes. pytest -q backends/transforms/test/test_replace_ops_with_channels_last_variants.py (23 passed). Authored with Codex.
1 parent 558473d commit ed7a23e

2 files changed

Lines changed: 26 additions & 1 deletion

File tree

‎backends/transforms/replace_ops_with_channels_last_variants.py‎

Lines changed: 6 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -133,6 +133,11 @@ class ReplaceOpsWithChannelsLastVariants(ExportPass):
133133
134134
By default, all currently implemented channels_last dialect ops are replaced.
135135
Pass a custom op_map to restrict or extend the set of replacements.
136+
137+
Metadata from each replaced operator is preserved so provenance and backend
138+
annotations survive the rewrite. ExportPass recomputes shape metadata after
139+
retracing. Callers must reject or remap semantic metadata tied to dimensions
140+
changed by ``input_indices`` or ``output_indices``, such as per-channel axes.
136141
"""
137142

138143
def __init__(
@@ -229,7 +234,7 @@ def call(self, graph_module: torch.fx.GraphModule) -> PassResult:
229234
args=tuple(args),
230235
kwargs=node.kwargs,
231236
)
232-
nhwc_node.meta = {}
237+
nhwc_node.meta = dict(node.meta)
233238

234239
users = list(node.users)
235240
if all(

‎backends/transforms/test/test_replace_ops_with_channels_last_variants.py‎

Lines changed: 20 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -132,6 +132,26 @@ def forward(self, x):
132132

133133

134134
class TestReplaceOpsWithChannelsLastVariants:
135+
def test_preserves_metadata_and_recomputes_shape(self):
136+
ep = _export_to_edge(Conv2dModule(), (torch.randn(1, 4, 8, 8),))
137+
conv = _find_nodes(ep.graph_module, exir_ops.edge.aten.convolution.default)[0]
138+
metadata = {
139+
"debug_handle": 1234,
140+
"from_node": [("source", "convolution")],
141+
"input_qparams": {0: "input"},
142+
"output_qparams": {0: "output"},
143+
}
144+
conv.meta.update(metadata)
145+
146+
result = ReplaceOpsWithChannelsLastVariants(ep)(ep.graph_module)
147+
replaced = _find_nodes(
148+
result.graph_module, exir_ops.edge.channels_last.convolution.default
149+
)[0]
150+
151+
for key, value in metadata.items():
152+
assert replaced.meta[key] == value
153+
assert tuple(replaced.meta["val"].shape) == (1, 8, 8, 4)
154+
135155
def test_conv2d(self):
136156
ep = _export_to_edge(Conv2dModule(bias=True), (torch.randn(1, 4, 8, 8),))
137157
assert _count(ep.graph_module, exir_ops.edge.aten.convolution.default) == 1

0 commit comments

Comments
 (0)