| 861 | return apply_reshape_impl(rdims); |
| 862 | } |
| 863 | bool 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()); |
nothing calls this directly
no test coverage detected