Skip to content

[Issue]: Try_compute_shape in eliminate_contiguous cannot compute the cumulative layout #5109

Description

@MQ-mengqing

Problem Description

When I wrote a new test with follows diff, I got terminate called after throwing an instance of 'std::runtime_error'\n what(): Not standard shape (at develop branch)

diff --git a/test/eliminate_contiguous_test.cpp b/test/eliminate_contiguous_test.cpp
index 78095b0ff..00a8babd4 100644
--- a/test/eliminate_contiguous_test.cpp
+++ b/test/eliminate_contiguous_test.cpp
@@ -27,6 +27,7 @@
 #include <migraphx/instruction.hpp>
 #include <basic_ops.hpp>
 #include <migraphx/make_op.hpp>
+#include <migraphx/op/pooling.hpp>
 
 #include <pointwise.hpp>
 #include <test.hpp>
@@ -289,4 +290,30 @@ TEST_CASE(slice_contiguous)
                m.begin(), m.end(), [](auto&& ins) { return ins.name() == "contiguous"; }) == 1);
 }
 
+TEST_CASE(chain_concat_nhwc_propagation)
+{
+    migraphx::module m;
+
+    auto s_nhwc =
+        migraphx::shape::from_permutation(migraphx::shape::float_type, {2, 3, 8, 8}, {0, 2, 3, 1});
+    auto d  = m.add_parameter("d", s_nhwc);
+    auto cd = m.add_instruction(migraphx::make_op("contiguous"), d);
+
+    auto pool_op = migraphx::make_op("pooling",
+                                     {{"mode", migraphx::op::pooling_mode::max},
+                                      {"padding", {0, 0}},
+                                      {"stride", {1, 1}},
+                                      {"lengths", {1, 1}}});
+    auto a       = m.add_instruction(pool_op, cd);
+    auto b       = m.add_instruction(pool_op, a);
+    auto c       = m.add_instruction(pool_op, b);
+    auto cc      = m.add_instruction(migraphx::make_op("concat", {{"axis", 1}}), a, b, c);
+    m.add_instruction(pass_standard_op{}, cc);
+
+    auto count = std::distance(m.begin(), m.end());
+    run_pass(m);
+    EXPECT(std::distance(m.begin(), m.end()) == count);
+    EXPECT(std::any_of(m.begin(), m.end(), [](auto&& ins) { return ins.name() == "contiguous"; }));
+}
+
 int main(int argc, const char* argv[]) { test::run(argc, argv); }

That's because try_compute_shape recursed per output and substituted only the current instruction's shape per branch, leaving every other input at its get_shape(). For consumers whose output layout is decided by find_permutation voting across multiple inputs (e.g. concat), the unchanged NCHW siblings outvoted the single NHWC change at every recursion step, so the function never observed the cumulative effect of eliminating a chain of contiguous. The replaced contiguous then surfaced NHWC at consumers that require standard layout.

Operating System

Ubuntu 24.04.4 LTS (Noble Numbat)

CPU

AMD Ryzen 7 5800H with Radeon Graphics

GPU

Other

Other

No response

ROCm Version

ROCm 6.0.0

Steps to Reproduce

No response

(Optional for Linux users) Output of /opt/rocm/bin/rocminfo --support

No response

Additional Information

No response

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions