已合并
【PR】: [fix] Concat问题修复 #1231
xchu42创建于 7月8日
【PR】: [fix] Concat问题修复 #1231
已合并
共 5 个文件变更+93-26
| @@ -308,6 +308,7 @@ inline __aicore__ void Concat16MultipleColumns(const ConcatParams<T, dim_size> & | |||
| 308 | TransposeParams trans_param = trans_para; | 308 | TransposeParams trans_param = trans_para; |
| 309 | trans_param.column_offset = column_align_cnt; | 309 | trans_param.column_offset = column_align_cnt; |
| 310 | trans_param.columns_cur_loop = col_not_align_cnt; | 310 | trans_param.columns_cur_loop = col_not_align_cnt; |
| 311 | + AscendC::PipeBarrier<PIPE_V>(); | ||
| 311 | FirstTransposeMatrix(tmp_buf1, *dst.tensor, trans_param, 0, dst.stride[0]); | 312 | FirstTransposeMatrix(tmp_buf1, *dst.tensor, trans_param, 0, dst.stride[0]); |
| 312 | trans_param.column_offset = trans_para.column_loop * max_columns; | 313 | trans_param.column_offset = trans_para.column_loop * max_columns; |
| 313 | trans_param.columns_cur_loop = trans_para.columns_cur_loop; | 314 | trans_param.columns_cur_loop = trans_para.columns_cur_loop; |
| @@ -316,7 +317,9 @@ inline __aicore__ void Concat16MultipleColumns(const ConcatParams<T, dim_size> & | |||
| 316 | uint32_t row_cnt = col_not_align_cnt + trans_para.columns_cur_loop; // 第二次转置前的总行数 | 317 | uint32_t row_cnt = col_not_align_cnt + trans_para.columns_cur_loop; // 第二次转置前的总行数 |
| 317 | constexpr uint32_t column_cnt = GetTotalColumns<T>(); // 第二次转置前的总列数 | 318 | constexpr uint32_t column_cnt = GetTotalColumns<T>(); // 第二次转置前的总列数 |
| 318 | LocalTensor<T> tmp_buf2 = tmp_buf1[column_cnt * row_cnt]; | 319 | LocalTensor<T> tmp_buf2 = tmp_buf1[column_cnt * row_cnt]; |
| 320 | + AscendC::PipeBarrier<PIPE_V>(); | ||
| 319 | SecondTranspose<T>(tmp_buf1, tmp_buf2, row_cnt, column_cnt); | 321 | SecondTranspose<T>(tmp_buf1, tmp_buf2, row_cnt, column_cnt); |
| 322 | + AscendC::PipeBarrier<PIPE_V>(); | ||
| 320 | 323 | ||
| 321 | // 3.拼接回dst | 324 | // 3.拼接回dst |
| 322 | struct RowParam row; | 325 | struct RowParam row; |
| @@ -393,6 +396,7 @@ inline __aicore__ void MultipleInputsConcat16Rows(const ConcatParams<T, dim_size | |||
| 393 | trans_param.row.rows_cur_loop = row.rows_cur_loop; | 396 | trans_param.row.rows_cur_loop = row.rows_cur_loop; |
| 394 | trans_param.column_offset = column_align_cnt; | 397 | trans_param.column_offset = column_align_cnt; |
| 395 | trans_param.columns_cur_loop = col_not_align_cnt; | 398 | trans_param.columns_cur_loop = col_not_align_cnt; |
| 399 | + AscendC::PipeBarrier<PIPE_V>(); | ||
| 396 | FirstTransposeMatrix(tmp_buf1, *dst.tensor, trans_param, 0, dst.stride[0]); | 400 | FirstTransposeMatrix(tmp_buf1, *dst.tensor, trans_param, 0, dst.stride[0]); |
| 397 | 401 | ||
| 398 | auto total_column_cnt = col_not_align_cnt; | 402 | auto total_column_cnt = col_not_align_cnt; |
| @@ -409,12 +413,13 @@ inline __aicore__ void MultipleInputsConcat16Rows(const ConcatParams<T, dim_size | |||
| 409 | sub_column_cnt += srcs[idx].shape[1]; | 413 | sub_column_cnt += srcs[idx].shape[1]; |
| 410 | total_column_cnt += srcs[idx].shape[1]; | 414 | total_column_cnt += srcs[idx].shape[1]; |
| 411 | } | 415 | } |
| 412 | - | 416 | + AscendC::PipeBarrier<PIPE_V>(); |
| 413 | // 2.第二次转置回来,转置策略为竖着取横着放,尽量增大repeat | 417 | // 2.第二次转置回来,转置策略为竖着取横着放,尽量增大repeat |
| 414 | constexpr uint32_t column_cnt = GetTotalColumns<T>(); | 418 | constexpr uint32_t column_cnt = GetTotalColumns<T>(); |
| 415 | uint32_t row_cnt = total_column_cnt; | 419 | uint32_t row_cnt = total_column_cnt; |
| 416 | LocalTensor<T> tmp_buf2 = tmp_buf1[column_cnt * row_cnt]; | 420 | LocalTensor<T> tmp_buf2 = tmp_buf1[column_cnt * row_cnt]; |
| 417 | SecondTranspose<T>(tmp_buf1, tmp_buf2, row_cnt, column_cnt); | 421 | SecondTranspose<T>(tmp_buf1, tmp_buf2, row_cnt, column_cnt); |
| 422 | + AscendC::PipeBarrier<PIPE_V>(); | ||
| 418 | 423 | ||
| 419 | // 3.拼接回dst | 424 | // 3.拼接回dst |
| 420 | DataCopyToDst<T, dim_size>(dst, tmp_buf2, row_cnt, row, column_align_cnt); | 425 | DataCopyToDst<T, dim_size>(dst, tmp_buf2, row_cnt, row, column_align_cnt); |
| @@ -495,6 +500,7 @@ inline __aicore__ void CopyThenTranspose16Rows(const ConcatParams<T, dim_size> & | |||
| 495 | 500 | ||
| 496 | // 由于该分支只会在dst上尾部column未对齐的时候进入,因此需要先处理dst上未对齐的部分 | 501 | // 由于该分支只会在dst上尾部column未对齐的时候进入,因此需要先处理dst上未对齐的部分 |
| 497 | // 将dst上未对齐的column拷贝到buf | 502 | // 将dst上未对齐的column拷贝到buf |
| 503 | + AscendC::PipeBarrier<PIPE_V>(); | ||
| 498 | CopyUnAlignDstToTmp<T, kMergedDimNum>(tmp_buf1, dst, row, column_align_cnt, tmp_stride0); | 504 | CopyUnAlignDstToTmp<T, kMergedDimNum>(tmp_buf1, dst, row, column_align_cnt, tmp_stride0); |
| 499 | AscendC::PipeBarrier<PIPE_V>(); | 505 | AscendC::PipeBarrier<PIPE_V>(); |
| 500 | curr_tmp_column += align; | 506 | curr_tmp_column += align; |
| @@ -58,23 +58,50 @@ Expression CalcForSmallTailKernel(AscNodeOutputs &node_outputs, uint32_t concat_ | |||
| 58 | return Symbol(buf_size); | 58 | return Symbol(buf_size); |
| 59 | } | 59 | } |
| 60 | 60 | ||
| 61 | -bool IsAllStaticAligned(AscNodeInputs &node_inputs, uint32_t concat_dim, int32_t align_size) { | 61 | +bool IsAllStaticAligned(const AscNode &node, int32_t align_size) { |
| 62 | - for (uint32_t i = 0; i < node_inputs.Size(); ++i) { | 62 | + AscNodeInputs node_inputs = node.inputs; |
| 63 | - auto axis = node_inputs[i].attr.repeats[concat_dim]; | 63 | + AscNodeOutputs node_outputs = node.outputs; |
| 64 | - for (uint32_t j = concat_dim + 1; j < node_inputs[i].attr.repeats.size(); ++j) { | 64 | + const auto &output_attr = node_outputs[0].attr; |
| 65 | - axis = sym::Mul(axis, node_inputs[i].attr.repeats[j]); | 65 | + const auto &input0_attr = node_inputs[0].attr; |
| 66 | - } | 66 | + GE_WARN_ASSERT(!output_attr.vectorized_axis.empty() && !output_attr.vectorized_strides.empty(), |
| 67 | + "vectorized_axis or vectorized_strides is empty, not aligned."); | ||
| 67 | 68 | ||
| 68 | - if (SymbolicUtils::StaticCheckEq(sym::Mod(axis, Symbol(align_size)), sym::kSymbolZero) != TriBool::kTrue) { | 69 | + size_t concat_dim = 0; |
| 69 | - GELOGD("The product of dims after concat_dim is %s, not aligned.", | 70 | + bool find_concat_dim = false; |
| 70 | - SymbolicUtils::ToString(sym::Mod(axis, Symbol(align_size))).c_str()); | 71 | + for (size_t i = 0; i < output_attr.vectorized_axis.size(); i++) { |
| 72 | + auto axis_id = output_attr.vectorized_axis[i]; | ||
| 73 | + auto it = std::find(output_attr.axis.begin(), output_attr.axis.end(), axis_id); | ||
| 74 | + GE_WARN_ASSERT(it != output_attr.axis.end(), "axis_id %ld not found in output[0].attr.axis, not aligned.", axis_id); | ||
| 75 | + const auto pos = static_cast<uint32_t>(std::distance(output_attr.axis.begin(), it)); | ||
| 76 | + if (SymbolicUtils::StaticCheckEq(input0_attr.repeats[pos], output_attr.repeats[pos]) != TriBool::kTrue) { | ||
| 77 | + concat_dim = i; | ||
| 78 | + find_concat_dim = true; | ||
| 79 | + GELOGD("find concat_dim: %zu in vectorized_axis.", concat_dim); | ||
| 80 | + break; | ||
| 81 | + } | ||
| 82 | + } | ||
| 83 | + | ||
| 84 | + GE_WARN_ASSERT(find_concat_dim, "not find concat dim in vectorized_axis, not aligned."); | ||
| 85 | + | ||
| 86 | + for (uint32_t i = 0; i < node_inputs.Size(); ++i) { | ||
| 87 | + const auto &input_attr = node_inputs[i].attr; | ||
| 88 | + auto axis_id = input_attr.vectorized_axis[concat_dim]; | ||
| 89 | + auto it = std::find(input_attr.axis.begin(), input_attr.axis.end(), axis_id); | ||
| 90 | + GE_WARN_ASSERT(it != input_attr.axis.end(), "axis_id %ld not found in input[%u].attr.axis, not aligned.", axis_id, | ||
| 91 | + i); | ||
| 92 | + const auto pos = static_cast<uint32_t>(std::distance(input_attr.axis.begin(), it)); | ||
| 93 | + const Expression axis_size = input_attr.repeats[pos] * output_attr.vectorized_strides[concat_dim]; | ||
| 94 | + if (SymbolicUtils::StaticCheckEq(sym::Mod(axis_size, Symbol(align_size)), sym::kSymbolZero) != TriBool::kTrue) { | ||
| 95 | + GELOGD("input[%u]: repeats[%u] * vectorized_strides[%zu] = %s, not aligned to %d.", i, pos, concat_dim, | ||
| 96 | + SymbolicUtils::ToString(sym::Mod(axis_size, Symbol(align_size))).c_str(), align_size); | ||
| 71 | return false; | 97 | return false; |
| 72 | } | 98 | } |
| 73 | } | 99 | } |
| 74 | return true; | 100 | return true; |
| 75 | } | 101 | } |
| 76 | 102 | ||
| 77 | -Expression CalcForDefaultKernel(AscNodeInputs &node_inputs, uint32_t concat_dim, bool flag) { | 103 | +Expression CalcForDefaultKernel(const AscNode &node, uint32_t concat_dim, bool flag) { |
| 104 | + AscNodeInputs node_inputs = node.inputs; | ||
| 78 | Expression max_axis_size = Symbol(0); | 105 | Expression max_axis_size = Symbol(0); |
| 79 | if (flag) { | 106 | if (flag) { |
| 80 | for (uint32_t i = 1; i < node_inputs.Size(); ++i) { | 107 | for (uint32_t i = 1; i < node_inputs.Size(); ++i) { |
| @@ -93,7 +120,7 @@ Expression CalcForDefaultKernel(AscNodeInputs &node_inputs, uint32_t concat_dim, | |||
| 93 | auto type_size = GetSizeByDataType(node_inputs[0].attr.dtype); | 120 | auto type_size = GetSizeByDataType(node_inputs[0].attr.dtype); |
| 94 | GE_ASSERT_TRUE(type_size != 0, "Invalid node inputs dtype, sizeof(T) = 0."); | 121 | GE_ASSERT_TRUE(type_size != 0, "Invalid node inputs dtype, sizeof(T) = 0."); |
| 95 | Expression min_tmp_buf_size = Symbol(0); | 122 | Expression min_tmp_buf_size = Symbol(0); |
| 96 | - bool is_aligned = IsAllStaticAligned(node_inputs, concat_dim, ALIGNSIZE32 / type_size); | 123 | + bool is_aligned = IsAllStaticAligned(node, ALIGNSIZE32 / type_size); |
| 97 | if (type_size == TYPESIZEEQ8) { | 124 | if (type_size == TYPESIZEEQ8) { |
| 98 | min_tmp_buf_size = | 125 | min_tmp_buf_size = |
| 99 | is_aligned ? Symbol(0) | 126 | is_aligned ? Symbol(0) |
| @@ -137,7 +164,7 @@ std::vector<std::unique_ptr<TmpBufDesc>> CalcConcatTmpSize(const AscNode &node) | |||
| 137 | bool concat_small_tail = false; | 164 | bool concat_small_tail = false; |
| 138 | (void)af::AttrUtils::GetBool(node.GetOpDesc(), "_concat_small_tail", concat_small_tail); | 165 | (void)af::AttrUtils::GetBool(node.GetOpDesc(), "_concat_small_tail", concat_small_tail); |
| 139 | const auto tmp_buf_size = concat_small_tail ? CalcForSmallTailKernel(node_outputs, concat_dim) | 166 | const auto tmp_buf_size = concat_small_tail ? CalcForSmallTailKernel(node_outputs, concat_dim) |
| 140 | - : CalcForDefaultKernel(node_inputs, concat_dim, flag); | 167 | + : CalcForDefaultKernel(node, concat_dim, flag); |
| 141 | if (SymbolicUtils::StaticCheckEq(tmp_buf_size, sym::kSymbolZero) == TriBool::kTrue) { | 168 | if (SymbolicUtils::StaticCheckEq(tmp_buf_size, sym::kSymbolZero) == TriBool::kTrue) { |
| 142 | GELOGI("%s does not require tmp buf", node.GetNamePtr()); | 169 | GELOGI("%s does not require tmp buf", node.GetNamePtr()); |
| 143 | return {}; | 170 | return {}; |
| @@ -156,22 +183,11 @@ std::vector<std::unique_ptr<TmpBufDesc>> CalcConcatTmpSize(const AscNode &node) | |||
| 156 | 183 | ||
| 157 | std::vector<std::unique_ptr<TmpBufDesc>> CalcConcatTmpSizeV2(const AscNode &node) { | 184 | std::vector<std::unique_ptr<TmpBufDesc>> CalcConcatTmpSizeV2(const AscNode &node) { |
| 158 | AscNodeInputs node_inputs = node.inputs; | 185 | AscNodeInputs node_inputs = node.inputs; |
| 159 | - AscNodeOutputs node_outputs = node.outputs; | ||
| 160 | GE_ASSERT_TRUE(node_inputs.Size() > 0); | 186 | GE_ASSERT_TRUE(node_inputs.Size() > 0); |
| 161 | - uint32_t concat_dim = 0; | ||
| 162 | - const auto num_dims = node_outputs[0].attr.repeats.size(); | ||
| 163 | - for (uint32_t idx = 0; idx < num_dims; ++idx) { | ||
| 164 | - const auto i = num_dims - idx - 1; | ||
| 165 | - if (node_outputs[0].attr.repeats[i] != node_inputs[0].attr.repeats[i]) { | ||
| 166 | - concat_dim = i; | ||
| 167 | - break; | ||
| 168 | - } | ||
| 169 | - } | ||
| 170 | auto type_size = GetSizeByDataType(node_inputs[0].attr.dtype); | 187 | auto type_size = GetSizeByDataType(node_inputs[0].attr.dtype); |
| 171 | GE_ASSERT_TRUE(type_size > 0, "%s Invalid node inputs dtype: %d", node.GetNamePtr(), | 188 | GE_ASSERT_TRUE(type_size > 0, "%s Invalid node inputs dtype: %d", node.GetNamePtr(), |
| 172 | static_cast<int32_t>(node_inputs[0].attr.dtype)); | 189 | static_cast<int32_t>(node_inputs[0].attr.dtype)); |
| 173 | - Expression min_tmp_buf_size = Symbol(0); | 190 | + bool is_aligned = IsAllStaticAligned(node, ALIGNSIZE32 / type_size); |
| 174 | - bool is_aligned = IsAllStaticAligned(node_inputs, concat_dim, ALIGNSIZE32 / type_size); | ||
| 175 | if (is_aligned) { | 191 | if (is_aligned) { |
| 176 | GELOGD("%s is all aligned", node.GetNamePtr()); | 192 | GELOGD("%s is all aligned", node.GetNamePtr()); |
| 177 | return {}; | 193 | return {}; |
| @@ -28,6 +28,7 @@ constexpr int32_t kConcatAlgTranspose = 0; | |||
| 28 | constexpr int64_t kVectorBlockSize = 256; | 28 | constexpr int64_t kVectorBlockSize = 256; |
| 29 | constexpr int64_t kMinGroupNum = 2L; | 29 | constexpr int64_t kMinGroupNum = 2L; |
| 30 | constexpr int64_t kMinGroupSizeByte = 1024U; | 30 | constexpr int64_t kMinGroupSizeByte = 1024U; |
| 31 | +constexpr af::float64_t kRatioThreshold = 2.0; | ||
| 31 | 32 | ||
| 32 | template <typename T1, typename T2> | 33 | template <typename T1, typename T2> |
| 33 | int64_t CeilDiv(const T1 n1, const T2 n2) { | 34 | int64_t CeilDiv(const T1 n1, const T2 n2) { |
| @@ -45,6 +46,7 @@ Status ConcatGroupPartitioner::Initialize() { | |||
| 45 | GE_ASSERT_NOTNULL(backend_spec); | 46 | GE_ASSERT_NOTNULL(backend_spec); |
| 46 | concat_by_transpose_ = (backend_spec->concat_alg == kConcatAlgTranspose); | 47 | concat_by_transpose_ = (backend_spec->concat_alg == kConcatAlgTranspose); |
| 47 | default_cols_per_group_ = kMaxBlockSize / dtype_size_; | 48 | default_cols_per_group_ = kMaxBlockSize / dtype_size_; |
| 49 | + max_cols_per_group_ = default_cols_per_group_; | ||
| 48 | max_input_num_per_group_ = MaxInputNumPerGroup(); | 50 | max_input_num_per_group_ = MaxInputNumPerGroup(); |
| 49 | // row足够大时仅row就足够分核, 不需要group parallel,否则尝试分组 | 51 | // row足够大时仅row就足够分核, 不需要group parallel,否则尝试分组 |
| 50 | if (known_rows_ <= kGroupParallelRowThreshold) { | 52 | if (known_rows_ <= kGroupParallelRowThreshold) { |
| @@ -56,6 +58,20 @@ Status ConcatGroupPartitioner::Initialize() { | |||
| 56 | const auto is_tail_concat = (concat_dim_ == output_attr.repeats.size() - 1UL); | 58 | const auto is_tail_concat = (concat_dim_ == output_attr.repeats.size() - 1UL); |
| 57 | can_use_small_tail_ = is_tail_concat && (dtype_size_ == sizeof(uint16_t) || dtype_size_ == sizeof(uint32_t)); | 59 | can_use_small_tail_ = is_tail_concat && (dtype_size_ == sizeof(uint16_t) || dtype_size_ == sizeof(uint32_t)); |
| 58 | group_type_to_limit_[kGroupTypeSmallTail] = kMaxBlockSizeForSmallTail; | 60 | group_type_to_limit_[kGroupTypeSmallTail] = kMaxBlockSizeForSmallTail; |
| 61 | + if (!is_tail_concat) { | ||
| 62 | + // 非尾轴concat, 尾轴会对齐到32B | ||
| 63 | + int64_t tail_dim_size = -1; | ||
| 64 | + GE_ASSERT_TRUE(concat_node_->outputs[0].attr.repeats.back().GetConstValue(tail_dim_size)); | ||
| 65 | + tail_dim_size *= dtype_size_; | ||
| 66 | + const auto aligned_size = (tail_dim_size + kAlignment - 1) / kAlignment * kAlignment; | ||
| 67 | + const auto ratio = static_cast<af::float64_t>(aligned_size) / static_cast<af::float64_t>(tail_dim_size); | ||
| 68 | + if (ratio >= kRatioThreshold) { | ||
| 69 | + default_cols_per_group_ = | ||
| 70 | + static_cast<int64_t>(static_cast<af::float64_t>(default_cols_per_group_) * kRatioThreshold / ratio); | ||
| 71 | + max_cols_per_group_ = default_cols_per_group_; | ||
| 72 | + GELOGD("tail dim = %ld, update default_cols_per_group to %ld", tail_dim_size, default_cols_per_group_); | ||
| 73 | + } | ||
| 74 | + } | ||
| 59 | } else { | 75 | } else { |
| 60 | if (output_cols_ > 0) { | 76 | if (output_cols_ > 0) { |
| 61 | GE_ASSERT_SUCCESS(TryOptimizeGroupSize()); | 77 | GE_ASSERT_SUCCESS(TryOptimizeGroupSize()); |
| @@ -232,7 +248,7 @@ bool ConcatGroupPartitioner::CanMerge(const ConcatGroupPartitioner::ConcatGroup | |||
| 232 | auto total_num = (lhs.end - lhs.start) + (rhs.end - rhs.start); | 248 | auto total_num = (lhs.end - lhs.start) + (rhs.end - rhs.start); |
| 233 | auto any_group_has_single_item = (lhs.end - lhs.start == 1) || (rhs.end - rhs.start == 1); | 249 | auto any_group_has_single_item = (lhs.end - lhs.start == 1) || (rhs.end - rhs.start == 1); |
| 234 | return (any_group_has_single_item || (total_num <= max_input_num_per_group_)) && | 250 | return (any_group_has_single_item || (total_num <= max_input_num_per_group_)) && |
| 235 | - ((lhs.size + rhs.size) <= kMaxBlockSize / dtype_size_); | 251 | + ((lhs.size + rhs.size) <= max_cols_per_group_); |
| 236 | } | 252 | } |
| 237 | 253 | ||
| 238 | void ConcatGroupPartitioner::ConvertToDefaultIfTooSmall() { | 254 | void ConcatGroupPartitioner::ConvertToDefaultIfTooSmall() { |
| @@ -99,6 +99,7 @@ class ConcatGroupPartitioner { | |||
| 99 | int64_t known_rows_ = 1L; | 99 | int64_t known_rows_ = 1L; |
| 100 | int64_t total_rows_ = 0L; | 100 | int64_t total_rows_ = 0L; |
| 101 | int64_t default_cols_per_group_ = 0L; | 101 | int64_t default_cols_per_group_ = 0L; |
| 102 | + int64_t max_cols_per_group_ = 0L; | ||
| 102 | bool has_recompute_ = false; | 103 | bool has_recompute_ = false; |
| 103 | }; | 104 | }; |
| 104 | } // namespace optimize | 105 | } // namespace optimize |
| @@ -125,12 +125,16 @@ void CreateStaticGraph(af::AscGraph &graph, int32_t dim0, int32_t dim1, int32_t | |||
| 125 | *x1.y.axis = {z0.id, zo_s_0.id}; | 125 | *x1.y.axis = {z0.id, zo_s_0.id}; |
| 126 | *x1.y.repeats = {s0, s1}; | 126 | *x1.y.repeats = {s0, s1}; |
| 127 | *x1.y.strides = {s1, One}; | 127 | *x1.y.strides = {s1, One}; |
| 128 | + *x1.y.vectorized_axis = {z0.id, zo_s_0.id}; | ||
| 129 | + *x1.y.vectorized_strides = {s1, One}; | ||
| 128 | 130 | ||
| 129 | x2.attr.sched.axis = {z0.id, zo_s_1.id}; | 131 | x2.attr.sched.axis = {z0.id, zo_s_1.id}; |
| 130 | x2.y.dtype = T; | 132 | x2.y.dtype = T; |
| 131 | *x2.y.axis = {z0.id, zo_s_1.id}; | 133 | *x2.y.axis = {z0.id, zo_s_1.id}; |
| 132 | *x2.y.repeats = {s0, s2}; | 134 | *x2.y.repeats = {s0, s2}; |
| 133 | *x2.y.strides = {s2, One}; | 135 | *x2.y.strides = {s2, One}; |
| 136 | + *x2.y.vectorized_axis = {z0.id, zo_s_1.id}; | ||
| 137 | + *x2.y.vectorized_strides = {s2, One}; | ||
| 134 | 138 | ||
| 135 | load1.x = x1.y; | 139 | load1.x = x1.y; |
| 136 | load1.attr.sched.axis = {z0.id, zo_s_0.id}; | 140 | load1.attr.sched.axis = {z0.id, zo_s_0.id}; |
| @@ -138,6 +142,8 @@ void CreateStaticGraph(af::AscGraph &graph, int32_t dim0, int32_t dim1, int32_t | |||
| 138 | *load1.y.axis = {z0.id, zo_s_0.id}; | 142 | *load1.y.axis = {z0.id, zo_s_0.id}; |
| 139 | *load1.y.repeats = {s0, s1}; | 143 | *load1.y.repeats = {s0, s1}; |
| 140 | *load1.y.strides = {s1, One}; | 144 | *load1.y.strides = {s1, One}; |
| 145 | + *load1.y.vectorized_axis = {z0.id, zo_s_0.id}; | ||
| 146 | + *load1.y.vectorized_strides = {s1, One}; | ||
| 141 | 147 | ||
| 142 | load2.x = x2.y; | 148 | load2.x = x2.y; |
| 143 | load2.attr.sched.axis = {z0.id, zo_s_1.id}; | 149 | load2.attr.sched.axis = {z0.id, zo_s_1.id}; |
| @@ -145,6 +151,8 @@ void CreateStaticGraph(af::AscGraph &graph, int32_t dim0, int32_t dim1, int32_t | |||
| 145 | *load2.y.axis = {z0.id, zo_s_1.id}; | 151 | *load2.y.axis = {z0.id, zo_s_1.id}; |
| 146 | *load2.y.repeats = {s0, s2}; | 152 | *load2.y.repeats = {s0, s2}; |
| 147 | *load2.y.strides = {s2, One}; | 153 | *load2.y.strides = {s2, One}; |
| 154 | + *load2.y.vectorized_axis = {z0.id, zo_s_1.id}; | ||
| 155 | + *load2.y.vectorized_strides = {s2, One}; | ||
| 148 | 156 | ||
| 149 | concat.x = {load1.y, load2.y}; | 157 | concat.x = {load1.y, load2.y}; |
| 150 | concat.attr.sched.axis = {z0.id, zo.id}; | 158 | concat.attr.sched.axis = {z0.id, zo.id}; |
| @@ -152,6 +160,8 @@ void CreateStaticGraph(af::AscGraph &graph, int32_t dim0, int32_t dim1, int32_t | |||
| 152 | *concat.y.axis = {z0.id, zo.id}; | 160 | *concat.y.axis = {z0.id, zo.id}; |
| 153 | *concat.y.repeats = {s0, s1 + s2}; | 161 | *concat.y.repeats = {s0, s1 + s2}; |
| 154 | *concat.y.strides = {s1 + s2, One}; | 162 | *concat.y.strides = {s1 + s2, One}; |
| 163 | + *concat.y.vectorized_axis = {z0.id, zo.id}; | ||
| 164 | + *concat.y.vectorized_strides = {s1 + s2, One}; | ||
| 155 | 165 | ||
| 156 | store.x = concat.y; | 166 | store.x = concat.y; |
| 157 | store.attr.sched.axis = {z0.id, zo.id}; | 167 | store.attr.sched.axis = {z0.id, zo.id}; |
| @@ -159,6 +169,8 @@ void CreateStaticGraph(af::AscGraph &graph, int32_t dim0, int32_t dim1, int32_t | |||
| 159 | *store.y.axis = {z0.id, zo.id}; | 169 | *store.y.axis = {z0.id, zo.id}; |
| 160 | *store.y.repeats = {s0, s1 + s2}; | 170 | *store.y.repeats = {s0, s1 + s2}; |
| 161 | *store.y.strides = {s1 + s2, One}; | 171 | *store.y.strides = {s1 + s2, One}; |
| 172 | + *store.y.vectorized_axis = {z0.id, zo.id}; | ||
| 173 | + *store.y.vectorized_strides = {s1 + s2, One}; | ||
| 162 | 174 | ||
| 163 | y.x = store.y; | 175 | y.x = store.y; |
| 164 | y.attr.sched.axis = {z0.id, zo.id}; | 176 | y.attr.sched.axis = {z0.id, zo.id}; |
| @@ -166,6 +178,8 @@ void CreateStaticGraph(af::AscGraph &graph, int32_t dim0, int32_t dim1, int32_t | |||
| 166 | *y.y.axis = {z0.id, zo.id}; | 178 | *y.y.axis = {z0.id, zo.id}; |
| 167 | *y.y.repeats = {s0, s1 + s2}; | 179 | *y.y.repeats = {s0, s1 + s2}; |
| 168 | *y.y.strides = {s1 + s2, One}; | 180 | *y.y.strides = {s1 + s2, One}; |
| 181 | + *y.y.vectorized_axis = {z0.id, zo.id}; | ||
| 182 | + *y.y.vectorized_strides = {s1 + s2, One}; | ||
| 169 | } | 183 | } |
| 170 | 184 | ||
| 171 | template <af::DataType T> | 185 | template <af::DataType T> |
| @@ -262,12 +276,16 @@ void CreateStaticGraphNotLastAxis(af::AscGraph &graph, int32_t dim0, int32_t dim | |||
| 262 | *x1.y.axis = {z0.id, zo_s_0.id}; | 276 | *x1.y.axis = {z0.id, zo_s_0.id}; |
| 263 | *x1.y.repeats = {s1, s0}; | 277 | *x1.y.repeats = {s1, s0}; |
| 264 | *x1.y.strides = {s0, One}; | 278 | *x1.y.strides = {s0, One}; |
| 279 | + *x1.y.vectorized_axis = {z0.id, zo_s_0.id}; | ||
| 280 | + *x1.y.vectorized_strides = {s0, One}; | ||
| 265 | 281 | ||
| 266 | x2.attr.sched.axis = {z0.id, zo_s_1.id}; | 282 | x2.attr.sched.axis = {z0.id, zo_s_1.id}; |
| 267 | x2.y.dtype = T; | 283 | x2.y.dtype = T; |
| 268 | *x2.y.axis = {z0.id, zo_s_1.id}; | 284 | *x2.y.axis = {z0.id, zo_s_1.id}; |
| 269 | *x2.y.repeats = {s2, s0}; | 285 | *x2.y.repeats = {s2, s0}; |
| 270 | *x2.y.strides = {s0, One}; | 286 | *x2.y.strides = {s0, One}; |
| 287 | + *x2.y.vectorized_axis = {z0.id, zo_s_1.id}; | ||
| 288 | + *x2.y.vectorized_strides = {s0, One}; | ||
| 271 | 289 | ||
| 272 | load1.x = x1.y; | 290 | load1.x = x1.y; |
| 273 | load1.attr.sched.axis = {z0.id, zo_s_0.id}; | 291 | load1.attr.sched.axis = {z0.id, zo_s_0.id}; |
| @@ -275,6 +293,8 @@ void CreateStaticGraphNotLastAxis(af::AscGraph &graph, int32_t dim0, int32_t dim | |||
| 275 | *load1.y.axis = {z0.id, zo_s_0.id}; | 293 | *load1.y.axis = {z0.id, zo_s_0.id}; |
| 276 | *load1.y.repeats = {s1, s0}; | 294 | *load1.y.repeats = {s1, s0}; |
| 277 | *load1.y.strides = {s0, One}; | 295 | *load1.y.strides = {s0, One}; |
| 296 | + *load1.y.vectorized_axis = {z0.id, zo_s_0.id}; | ||
| 297 | + *load1.y.vectorized_strides = {s0, One}; | ||
| 278 | 298 | ||
| 279 | load2.x = x2.y; | 299 | load2.x = x2.y; |
| 280 | load2.attr.sched.axis = {z0.id, zo_s_1.id}; | 300 | load2.attr.sched.axis = {z0.id, zo_s_1.id}; |
| @@ -282,6 +302,8 @@ void CreateStaticGraphNotLastAxis(af::AscGraph &graph, int32_t dim0, int32_t dim | |||
| 282 | *load2.y.axis = {z0.id, zo_s_1.id}; | 302 | *load2.y.axis = {z0.id, zo_s_1.id}; |
| 283 | *load2.y.repeats = {s2, s0}; | 303 | *load2.y.repeats = {s2, s0}; |
| 284 | *load2.y.strides = {s0, One}; | 304 | *load2.y.strides = {s0, One}; |
| 305 | + *load2.y.vectorized_axis = {z0.id, zo_s_1.id}; | ||
| 306 | + *load2.y.vectorized_strides = {s0, One}; | ||
| 285 | 307 | ||
| 286 | concat.x = {load1.y, load2.y}; | 308 | concat.x = {load1.y, load2.y}; |
| 287 | concat.attr.sched.axis = {z0.id, zo.id}; | 309 | concat.attr.sched.axis = {z0.id, zo.id}; |
| @@ -289,6 +311,8 @@ void CreateStaticGraphNotLastAxis(af::AscGraph &graph, int32_t dim0, int32_t dim | |||
| 289 | *concat.y.axis = {z0.id, zo.id}; | 311 | *concat.y.axis = {z0.id, zo.id}; |
| 290 | *concat.y.repeats = {s1 + s2, s0}; | 312 | *concat.y.repeats = {s1 + s2, s0}; |
| 291 | *concat.y.strides = {s0, One}; | 313 | *concat.y.strides = {s0, One}; |
| 314 | + *concat.y.vectorized_axis = {z0.id, zo.id}; | ||
| 315 | + *concat.y.vectorized_strides = {s0, One}; | ||
| 292 | 316 | ||
| 293 | store.x = concat.y; | 317 | store.x = concat.y; |
| 294 | store.attr.sched.axis = {z0.id, zo.id}; | 318 | store.attr.sched.axis = {z0.id, zo.id}; |
| @@ -297,6 +321,8 @@ void CreateStaticGraphNotLastAxis(af::AscGraph &graph, int32_t dim0, int32_t dim | |||
| 297 | *store.y.axis = {z0.id, zo.id}; | 321 | *store.y.axis = {z0.id, zo.id}; |
| 298 | *store.y.repeats = {s1 + s2, s0}; | 322 | *store.y.repeats = {s1 + s2, s0}; |
| 299 | *store.y.strides = {s0, One}; | 323 | *store.y.strides = {s0, One}; |
| 324 | + *store.y.vectorized_axis = {z0.id, zo.id}; | ||
| 325 | + *store.y.vectorized_strides = {s0, One}; | ||
| 300 | 326 | ||
| 301 | y.x = store.y; | 327 | y.x = store.y; |
| 302 | y.attr.sched.axis = {z0.id, zo.id}; | 328 | y.attr.sched.axis = {z0.id, zo.id}; |
| @@ -304,6 +330,8 @@ void CreateStaticGraphNotLastAxis(af::AscGraph &graph, int32_t dim0, int32_t dim | |||
| 304 | *y.y.axis = {z0.id, zo.id}; | 330 | *y.y.axis = {z0.id, zo.id}; |
| 305 | *y.y.repeats = {s1 + s2, s0}; | 331 | *y.y.repeats = {s1 + s2, s0}; |
| 306 | *y.y.strides = {s0, One}; | 332 | *y.y.strides = {s0, One}; |
| 333 | + *y.y.vectorized_axis = {z0.id, zo.id}; | ||
| 334 | + *y.y.vectorized_strides = {s0, One}; | ||
| 307 | } | 335 | } |
| 308 | 336 | ||
| 309 | /** | 337 | /** |