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

Function generate_shape_transforms_for

src/shape_transform_descriptor.cpp:2111–2236  ·  view source on GitHub ↗

Generate the shape transforms for strided view

Source from the content-addressed store, hash-verified

2109
2110// Generate the shape transforms for strided view
2111optional<std::vector<operation>>
2112generate_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(

Callers 2

generate_forFunction · 0.85
transform_indicesMethod · 0.85

Calls 15

accumulateFunction · 0.85
adjust_strided_shapeFunction · 0.85
unbroadcastFunction · 0.85
select_maskFunction · 0.85
lensMethod · 0.80
frontMethod · 0.80
make_opFunction · 0.70
find_permutationFunction · 0.70
invert_permutationFunction · 0.70
reorder_shapeFunction · 0.70
any_ofFunction · 0.50
transformFunction · 0.50

Tested by

no test coverage detected