|
13 | 13 | TosaPipelineFP, |
14 | 14 | TosaPipelineINT, |
15 | 15 | ) |
| 16 | +from executorch.backends.test.harness.stages import StageType |
16 | 17 |
|
17 | 18 |
|
18 | 19 | InputT = Tuple[Any, ...] |
@@ -540,15 +541,6 @@ def forward(self, x: torch.Tensor): |
540 | 541 | (torch.randn(1, 2, 8, 8).to(memory_format=torch.channels_last),), |
541 | 542 | 3, |
542 | 543 | ), |
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 | | - ), |
552 | 544 | "cumsum_rank4_dim3_channels_last": TransposeCountCase( |
553 | 545 | CumsumModule(), |
554 | 546 | (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 |
582 | 574 | pipeline = TosaPipelineFP[InputT](case.module, case.inputs, aten_op=[]) |
583 | 575 | pipeline.count_tosa_ops({"TRANSPOSE": case.expected_transposes}) |
584 | 576 | 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