set reduce_dim, left_dim and update x_dim eg: x_dim = [2, 4, 6] origin_reduce_dims = [0, 1] --SetReduceDim--> x_dim = [8,6], reduce_dim = [0], left_dim = [1]
| 301 | // eg: x_dim = [2, 4, 6] origin_reduce_dims = [0, 1] |
| 302 | // --SetReduceDim--> x_dim = [8,6], reduce_dim = [0], left_dim = [1] |
| 303 | void SetReduceDim() { |
| 304 | std::set<int64_t> reduce_set; |
| 305 | for (auto e : reduce_dims_origin) { |
| 306 | auto pos = e >= 0 ? e : e + x_dim.size(); |
| 307 | reduce_set.insert(pos); |
| 308 | } |
| 309 | |
| 310 | std::vector<int64_t> reduce_dim_temp(reduce_set.begin(), reduce_set.end()); |
| 311 | std::sort(reduce_dim_temp.begin(), reduce_dim_temp.end()); |
| 312 | |
| 313 | // update reduce_dim and x_dim |
| 314 | std::vector<int64_t> x_new_dim; |
| 315 | |
| 316 | reduce_dim.push_back(reduce_dim_temp[0]); |
| 317 | x_new_dim.push_back(x_dim[0]); |
| 318 | |
| 319 | int64_t idx_reduce = 1; |
| 320 | int64_t num = 0; |
| 321 | |
| 322 | if (reduce_dim_temp.size() > 1) { |
| 323 | for (int64_t i = 1; i < x_dim.size(); i++) { |
| 324 | if ((idx_reduce < reduce_dim_temp.size()) && |
| 325 | (i == reduce_dim_temp[idx_reduce])) { |
| 326 | int64_t result = |
| 327 | reduce_dim_temp[idx_reduce] - reduce_dim[reduce_dim.size() - 1]; |
| 328 | bool is_equal = ((result - num) == 1); |
| 329 | if (is_equal) { |
| 330 | x_new_dim[x_new_dim.size() - 1] *= x_dim[i]; |
| 331 | num++; |
| 332 | } else { |
| 333 | reduce_dim.push_back(reduce_dim_temp[idx_reduce] - num); |
| 334 | x_new_dim.push_back(x_dim[i]); |
| 335 | } |
| 336 | idx_reduce++; |
| 337 | } else { |
| 338 | x_new_dim.push_back(x_dim[i]); |
| 339 | } |
| 340 | } |
| 341 | } else { |
| 342 | x_new_dim = x_dim; |
| 343 | } |
| 344 | |
| 345 | // update x_dim |
| 346 | x_dim = x_new_dim; |
| 347 | std::vector<int64_t>().swap(x_new_dim); |
| 348 | |
| 349 | std::vector<int64_t> reduce_dim_new; |
| 350 | int64_t is_reduced = 0; |
| 351 | for (auto e : reduce_dim) { |
| 352 | is_reduced |= 1 << e; |
| 353 | } |
| 354 | |
| 355 | std::vector<int64_t>().swap(reduce_dim); |
| 356 | |
| 357 | for (int64_t i = 0; i < x_dim.size(); i++) { |
| 358 | if ((i == 0) || (((is_reduced >> i) ^ (is_reduced >> (i - 1))) & 1)) { |
| 359 | x_new_dim.push_back(x_dim[i]); |
| 360 | if ((is_reduced >> i) & 1) |