Generate the shape transforms for strided view
| 2109 | |
| 2110 | // Generate the shape transforms for strided view |
| 2111 | optional<std::vector<operation>> |
| 2112 | generate_shape_transforms_for(shape s, const std::vector<std::size_t>& idims, std::int64_t offset) |
| 2113 | { |
| 2114 | std::vector<operation> result; |
| 2115 | if(s.lens().empty()) |
| 2116 | return std::nullopt; |
| 2117 | std::size_t ielements = |
| 2118 | std::accumulate(idims.begin(), idims.end(), std::size_t(1), std::multiplies<>()); |
| 2119 | auto extra = adjust_strided_shape(s, ielements); |
| 2120 | // TODO: Improve handling of multiple dimensions, for now just reshape to 1 dimension |
| 2121 | if(idims.size() != 1) |
| 2122 | { |
| 2123 | result.push_back(make_op("reshape", {{"dims", {ielements}}})); |
| 2124 | auto ops = generate_shape_transforms_for(s, {ielements}, offset); |
| 2125 | if(not ops) |
| 2126 | return std::nullopt; |
| 2127 | result.insert(result.end(), ops->begin(), ops->end()); |
| 2128 | return result; |
| 2129 | } |
| 2130 | auto pre_broadcast = unbroadcast(s); |
| 2131 | auto perm = find_permutation(pre_broadcast); |
| 2132 | auto iperm = invert_permutation(perm); |
| 2133 | auto pre_transpose = reorder_shape(pre_broadcast, perm); |
| 2134 | |
| 2135 | std::vector<std::size_t> start_lens; |
| 2136 | std::adjacent_difference(pre_transpose.strides().begin(), |
| 2137 | pre_transpose.strides().end(), |
| 2138 | std::back_inserter(start_lens), |
| 2139 | [](auto y, auto x) -> std::size_t { |
| 2140 | assert(x >= y); |
| 2141 | assert(y != 0); |
| 2142 | if((x % y) != 0) |
| 2143 | return 0; |
| 2144 | return x / y; |
| 2145 | }); |
| 2146 | if(std::any_of(start_lens.begin(), start_lens.end(), [](auto len) { return len == 0; })) |
| 2147 | return std::nullopt; |
| 2148 | start_lens.front() = extra > 1 ? extra : pre_transpose.lens().front(); |
| 2149 | |
| 2150 | std::size_t nelements = |
| 2151 | std::accumulate(start_lens.begin(), start_lens.end(), std::size_t(1), std::multiplies<>()); |
| 2152 | |
| 2153 | if(nelements < pre_transpose.elements() * extra) |
| 2154 | return std::nullopt; |
| 2155 | |
| 2156 | std::vector<std::size_t> start_mask(start_lens.size(), 0); |
| 2157 | if(offset != 0) |
| 2158 | { |
| 2159 | shape start_shape{shape::float_type, start_lens}; |
| 2160 | auto idx = start_shape.multi(offset); |
| 2161 | |
| 2162 | std::vector<std::size_t> overhead; |
| 2163 | std::transform(start_lens.begin(), |
| 2164 | start_lens.end(), |
| 2165 | pre_transpose.lens().begin(), |
| 2166 | std::back_inserter(overhead), |
| 2167 | [](auto start_len, auto len) { return start_len - len; }); |
| 2168 | if(std::equal( |
no test coverage detected