MCPcopy Create free account
hub / github.com/ROCm/AMDMIGraphX / apply_reshape_impl

Method apply_reshape_impl

src/shape_transform_descriptor.cpp:863–938  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

861 return apply_reshape_impl(rdims);
862}
863bool shape_transform_descriptor::apply_reshape_impl(const std::vector<std::size_t>& rdims)
864{
865 assert(migraphx::elements(rdims) == this->elements());
866 if(migraphx::equal(
867 dimensions, rdims, [](const dimension& d, std::size_t rdim) { return d.len() == rdim; }))
868 return true;
869 std::vector<dimension> new_dims;
870 auto subs = get_all_subdimensions(dimensions);
871 std::size_t i = 0;
872 std::size_t r = 0;
873 while(i < subs.size() and r < rdims.size())
874 {
875 const auto& sub = subs[i];
876 auto idim = sub.len;
877 auto rdim = rdims[r];
878 if(idim == rdim)
879 {
880 new_dims.push_back({{sub}});
881 }
882 // squeeze
883 else if(rdim > idim)
884 {
885 auto start = subs.begin() + i;
886 auto it = compute_end_dim(start, subs.end(), rdim, std::mem_fn(&dimension::sub::len));
887 if(it == start)
888 return false;
889 assert(it != subs.end());
890 auto n = it - start;
891 i += n;
892 new_dims.push_back({{start, it + 1}});
893 }
894 // unsqueeze
895 else // if(rdim < idim)
896 {
897 auto start = rdims.begin() + r;
898 auto it = compute_end_dim(start, rdims.end(), idim, id{});
899 if(it == start)
900 return false;
901 assert(it != rdims.end());
902 auto n = it - start;
903 r += n;
904 transform(range(n + 1), std::back_inserter(new_dims), [&](auto j) -> dimension {
905 auto new_sub = sub;
906 new_sub.add_split_axis(j);
907 new_sub.len = start[j];
908 return {{new_sub}};
909 });
910 }
911 r++;
912 i++;
913 }
914
915 // Handle trailing 1s
916 if(new_dims.size() < rdims.size() and not new_dims.empty())
917 {
918 auto* sub = get_last_subdimension(new_dims);
919 auto axis = sub == nullptr ? std::vector<std::size_t>{} : sub->axis;
920 auto trailing_dims = range(rdims.begin() + new_dims.size(), rdims.end());

Callers

nothing calls this directly

Calls 15

elementsMethod · 0.95
get_all_subdimensionsFunction · 0.85
get_last_subdimensionFunction · 0.85
distanceFunction · 0.85
lenMethod · 0.80
add_split_axisMethod · 0.80
compute_end_dimFunction · 0.70
elementsFunction · 0.50
equalFunction · 0.50
transformFunction · 0.50
rangeFunction · 0.50
any_ofFunction · 0.50

Tested by

no test coverage detected