已合并
【PR】: [fix] Concat问题修复 #1231
【PR】: [fix] Concat问题修复 #1231
已合并
xchu42创建于 7月8日
共 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.拼接回dst324 // 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.第二次转置回来,转置策略为竖着取横着放,尽量增大repeat417 // 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.拼接回dst424 // 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拷贝到buf502 // 将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 
157std::vector<std::unique_ptr<TmpBufDesc>> CalcConcatTmpSizeV2(const AscNode &node) {184std::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;
28constexpr int64_t kVectorBlockSize = 256;28constexpr int64_t kVectorBlockSize = 256;
29constexpr int64_t kMinGroupNum = 2L;29constexpr int64_t kMinGroupNum = 2L;
30constexpr int64_t kMinGroupSizeByte = 1024U;30constexpr int64_t kMinGroupSizeByte = 1024U;
31+constexpr af::float64_t kRatioThreshold = 2.0;
31 32 
32template <typename T1, typename T2>33template <typename T1, typename T2>
33int64_t CeilDiv(const T1 n1, const T2 n2) {34int64_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 
238void ConcatGroupPartitioner::ConvertToDefaultIfTooSmall() {254void 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 optimize105} // 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 
171template <af::DataType T>185template <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/**