Skip to content

Commit 89a8e59

Browse files
Arm backend: Stabilize GroupNorm transpose-count
Preserve the unmodified GroupNorm model and select the exact expected transpose count from its exported output layout before backend lowering. This supports release and nightly PyTorch while retaining native GroupNorm regression coverage and numerical comparisons. Authored with assistance from OpenAI Codex. Signed-off-by: Sangwon Ha <sangwon.ha@arm.com> Change-Id: I25aa7f4269fe9676579b43b57119850a8ebac45b
1 parent bb2683b commit 89a8e59

1 file changed

Lines changed: 22 additions & 9 deletions

File tree

‎backends/arm/test/misc/test_transpose_counts.py‎

Lines changed: 22 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -13,6 +13,7 @@
1313
TosaPipelineFP,
1414
TosaPipelineINT,
1515
)
16+
from executorch.backends.test.harness.stages import StageType
1617

1718

1819
InputT = Tuple[Any, ...]
@@ -540,15 +541,6 @@ def forward(self, x: torch.Tensor):
540541
(torch.randn(1, 2, 8, 8).to(memory_format=torch.channels_last),),
541542
3,
542543
),
543-
# Rebaselined 1 -> 2 with the dev20260913 pin, which changed how GroupNorm
544-
# decomposes under channels-last. The pipeline still compares against the
545-
# TOSA reference model, so the extra transpose is numerically correct, but
546-
# nobody has checked whether it is one the backend could still fold away.
547-
"groupnorm_channels_last": TransposeCountCase(
548-
GroupNormModule(),
549-
(torch.randn(1, 4, 4, 4).to(memory_format=torch.channels_last),),
550-
2,
551-
),
552544
"cumsum_rank4_dim3_channels_last": TransposeCountCase(
553545
CumsumModule(),
554546
(torch.randn(1, 2, 3, 4).to(memory_format=torch.channels_last), 3),
@@ -582,3 +574,24 @@ def test_transpose_counts_tosa_FP_channels_last(case: TransposeCountCase) -> Non
582574
pipeline = TosaPipelineFP[InputT](case.module, case.inputs, aten_op=[])
583575
pipeline.count_tosa_ops({"TRANSPOSE": case.expected_transposes})
584576
pipeline.run()
577+
578+
579+
def test_transpose_counts_tosa_FP_groupnorm_channels_last() -> None:
580+
inputs = (torch.randn(1, 4, 4, 4).to(memory_format=torch.channels_last),)
581+
pipeline = TosaPipelineFP[InputT](
582+
GroupNormModule(), inputs, aten_op="torch.ops.aten.group_norm.default"
583+
)
584+
pipeline.tester.export()
585+
pipeline.pop_stage("export")
586+
587+
# PyTorch's output layout determines the boundary transposes before lowering.
588+
exported_program = pipeline.tester.get_artifact(StageType.EXPORT)
589+
output_node = next(n for n in exported_program.graph.nodes if n.op == "output")
590+
(output,) = output_node.args[0]
591+
expected_transposes = {
592+
(0, 1, 2, 3): 1,
593+
(0, 2, 3, 1): 2,
594+
}[output.meta["val"].dim_order()]
595+
596+
pipeline.count_tosa_ops({"TRANSPOSE": expected_transposes})
597+
pipeline.run()

0 commit comments

Comments
 (0)