已合并
fix: 完善 IndirectLoad SIMD/SIMT codegen 与回归测试 #1828
xiebangrui2025创建于 8月25日
fix: 完善 IndirectLoad SIMD/SIMT codegen 与回归测试 #1828
已合并
共 16 个文件变更+2548-235
| @@ -12,6 +12,7 @@ | |||
| 12 | 12 | ||
| 13 | 13 | ||
| 14 | 14 | ||
| 15 | + | ||
| 15 | 16 | ||
| 16 | 17 | ||
| 17 | 18 | ||
| @@ -52,6 +53,11 @@ bool TryClassifyDynamicShapeLayout(const LogicalTensorView &logical, IndirectLoa | |||
| 52 | if (!has_dynamic_shape) { | 53 | if (!has_dynamic_shape) { |
| 53 | return false; | 54 | return false; |
| 54 | } | 55 | } |
| 56 | + // A dynamic outer dimension does not prevent proving a compact zero-stride | ||
| 57 | + // view when all non-broadcast dimensions are contiguous. Preserve this | ||
| 58 | + // producer-side layout so Broadcast and its source use the same tensor view. | ||
| 59 | + bool has_zero_stride = false; | ||
| 60 | + af::Expression physical_span = af::sym::kSymbolOne; | ||
| 55 | for (size_t dim = 0UL; dim < logical.sizes.size(); ++dim) { | 61 | for (size_t dim = 0UL; dim < logical.sizes.size(); ++dim) { |
| 56 | const auto dim_kind = ClassifyTensorDim(logical.sizes[dim], logical.strides[dim]); | 62 | const auto dim_kind = ClassifyTensorDim(logical.sizes[dim], logical.strides[dim]); |
| 57 | if (dim_kind == TensorDimKind::kIllegal) { | 63 | if (dim_kind == TensorDimKind::kIllegal) { |
| @@ -59,8 +65,25 @@ bool TryClassifyDynamicShapeLayout(const LogicalTensorView &logical, IndirectLoa | |||
| 59 | } | 65 | } |
| 60 | if (dim_kind == TensorDimKind::kZeroStride) { | 66 | if (dim_kind == TensorDimKind::kZeroStride) { |
| 61 | layout.physical_repeats[dim] = af::sym::kSymbolOne; | 67 | layout.physical_repeats[dim] = af::sym::kSymbolOne; |
| 68 | + has_zero_stride = true; | ||
| 62 | } | 69 | } |
| 63 | } | 70 | } |
| 71 | + if (has_zero_stride) { | ||
| 72 | + physical_span = af::sym::kSymbolOne; | ||
| 73 | + for (size_t index = logical.sizes.size(); index > 0UL; --index) { | ||
| 74 | + const size_t dim = index - 1UL; | ||
| 75 | + if (ClassifyTensorDim(logical.sizes[dim], logical.strides[dim]) == TensorDimKind::kZeroStride) { | ||
| 76 | + continue; | ||
| 77 | + } | ||
| 78 | + if (af::SymbolicUtils::StaticCheckEq(logical.strides[dim], physical_span) != af::TriBool::kTrue) { | ||
| 79 | + layout.kind = IndirectLoadLayoutKind::kStrided; | ||
| 80 | + return true; | ||
| 81 | + } | ||
| 82 | + physical_span = physical_span + (logical.sizes[dim] - af::sym::kSymbolOne) * logical.strides[dim]; | ||
| 83 | + } | ||
| 84 | + layout.kind = IndirectLoadLayoutKind::kZeroStrideCompact; | ||
| 85 | + return true; | ||
| 86 | + } | ||
| 64 | layout.kind = IndirectLoadLayoutKind::kStrided; | 87 | layout.kind = IndirectLoadLayoutKind::kStrided; |
| 65 | return true; | 88 | return true; |
| 66 | } | 89 | } |
| @@ -111,6 +134,12 @@ TemplateBehavior GetBehavior(TemplateRole role) { | |||
| 111 | behavior.uses_direct_gm_pipeline = true; | 134 | behavior.uses_direct_gm_pipeline = true; |
| 112 | behavior.preserves_vectorized_axis = true; | 135 | behavior.preserves_vectorized_axis = true; |
| 113 | break; | 136 | break; |
| 137 | + case TemplateRole::kSimtFanoutBranch: | ||
| 138 | + behavior.excludes_tiling_group = true; | ||
| 139 | + behavior.skips_main_schedule_tiling = true; | ||
| 140 | + behavior.uses_direct_gm_pipeline = true; | ||
| 141 | + behavior.preserves_vectorized_axis = true; | ||
| 142 | + break; | ||
| 114 | case TemplateRole::kSimtOp: | 143 | case TemplateRole::kSimtOp: |
| 115 | behavior.uses_direct_gm_pipeline = true; | 144 | behavior.uses_direct_gm_pipeline = true; |
| 116 | behavior.skips_ub_lifecycle = true; | 145 | behavior.skips_ub_lifecycle = true; |
| @@ -158,18 +187,39 @@ struct PostReduceChain { | |||
| 158 | }; | 187 | }; |
| 159 | 188 | ||
| 160 | PostReduceChain FindPostReduceChain(const af::AscNodePtr &node) { | 189 | PostReduceChain FindPostReduceChain(const af::AscNodePtr &node) { |
| 161 | - af::AscNodePtr producer = node; | 190 | + if (node == nullptr) { |
| 162 | - while (producer != nullptr) { | 191 | + return {}; |
| 163 | - const af::AscNodePtr consumer = GetOnlyOutputConsumer(producer); | ||
| 164 | - if (consumer == nullptr) { | ||
| 165 | - return {}; | ||
| 166 | - } | ||
| 167 | - if (consumer->attr.api.compute_type == af::ComputeType::kComputeReduce) { | ||
| 168 | - return {producer, consumer}; | ||
| 169 | - } | ||
| 170 | - producer = consumer; | ||
| 171 | } | 192 | } |
| 172 | - return {}; | 193 | + std::vector<af::AscNodePtr> pending; |
| 194 | + std::unordered_set<af::AscNode *> visited; | ||
| 195 | + for (const auto &out_node : node->GetOutDataNodes()) { | ||
| 196 | + const auto out_asc_node = std::dynamic_pointer_cast<af::AscNode>(out_node); | ||
| 197 | + if (out_asc_node != nullptr && visited.emplace(out_asc_node.get()).second) { | ||
| 198 | + pending.emplace_back(out_asc_node); | ||
| 199 | + } | ||
| 200 | + } | ||
| 201 | + | ||
| 202 | + af::AscNodePtr reduce; | ||
| 203 | + for (size_t index = 0UL; index < pending.size(); ++index) { | ||
| 204 | + const auto ¤t = pending[index]; | ||
| 205 | + if (current->attr.api.compute_type == af::ComputeType::kComputeReduce) { | ||
| 206 | + if (reduce != nullptr && reduce != current) { | ||
| 207 | + return {}; | ||
| 208 | + } | ||
| 209 | + reduce = current; | ||
| 210 | + continue; | ||
| 211 | + } | ||
| 212 | + for (const auto &out_node : current->GetOutDataNodes()) { | ||
| 213 | + const auto out_asc_node = std::dynamic_pointer_cast<af::AscNode>(out_node); | ||
| 214 | + if (out_asc_node != nullptr && visited.emplace(out_asc_node.get()).second) { | ||
| 215 | + pending.emplace_back(out_asc_node); | ||
| 216 | + } | ||
| 217 | + } | ||
| 218 | + } | ||
| 219 | + if (reduce == nullptr) { | ||
| 220 | + return {}; | ||
| 221 | + } | ||
| 222 | + return {GetInputProducer(reduce, 0UL), reduce}; | ||
| 173 | } | 223 | } |
| 174 | } // namespace | 224 | } // namespace |
| 175 | 225 | ||
| @@ -182,6 +232,38 @@ af::AscNodePtr GetPostReduceInputProducer(const af::AscNodePtr &node) { | |||
| 182 | } | 232 | } |
| 183 | 233 | ||
| 184 | bool ShouldSkipTpipeTensorCollection(const af::AscNodePtr &node) { | 234 | bool ShouldSkipTpipeTensorCollection(const af::AscNodePtr &node) { |
| 235 | + // A SIMT IndirectLoad normally bypasses the TPipe because its direct-GM | ||
| 236 | + // chain consumes the value in registers. With a user fan-out, however, | ||
| 237 | + // an ordinary scheduled branch may consume the IndirectLoad result as a | ||
| 238 | + // UB tensor (for example through a VectorFunc), so keep that output in the | ||
| 239 | + // tensor table for the branch while retaining the direct-GM path. | ||
| 240 | + if (node != nullptr && af::ops::IsOps<af::ascir_op::IndirectLoad>(node)) { | ||
| 241 | + std::vector<af::AscNodePtr> pending; | ||
| 242 | + std::unordered_set<const af::AscNode *> visited; | ||
| 243 | + for (const auto &out : node->GetOutDataNodes()) { | ||
| 244 | + const auto consumer = std::dynamic_pointer_cast<af::AscNode>(out); | ||
| 245 | + if (consumer != nullptr && visited.emplace(consumer.get()).second) { | ||
| 246 | + pending.emplace_back(consumer); | ||
| 247 | + } | ||
| 248 | + } | ||
| 249 | + size_t store_count = 0UL; | ||
| 250 | + for (size_t i = 0UL; i < pending.size(); ++i) { | ||
| 251 | + const auto ¤t = pending[i]; | ||
| 252 | + if (af::ops::IsOps<af::ascir_op::Store>(current)) { | ||
| 253 | + ++store_count; | ||
| 254 | + if (store_count > 1UL) { | ||
| 255 | + return false; | ||
| 256 | + } | ||
| 257 | + continue; | ||
| 258 | + } | ||
| 259 | + for (const auto &out : current->GetOutDataNodes()) { | ||
| 260 | + const auto consumer = std::dynamic_pointer_cast<af::AscNode>(out); | ||
| 261 | + if (consumer != nullptr && visited.emplace(consumer.get()).second) { | ||
| 262 | + pending.emplace_back(consumer); | ||
| 263 | + } | ||
| 264 | + } | ||
| 265 | + } | ||
| 266 | + } | ||
| 185 | const TemplateBehavior behavior = GetTemplateBehavior(node); | 267 | const TemplateBehavior behavior = GetTemplateBehavior(node); |
| 186 | const af::AscNodePtr consumer = GetOnlyOutputConsumer(node); | 268 | const af::AscNodePtr consumer = GetOnlyOutputConsumer(node); |
| 187 | return (behavior.skips_api_emit || behavior.skips_ub_lifecycle) && | 269 | return (behavior.skips_api_emit || behavior.skips_ub_lifecycle) && |
| @@ -258,7 +340,8 @@ af::Status GetTemplateLogicalView(const af::AscNodePtr &node, TemplateLogicalVie | |||
| 258 | return af::SUCCESS; | 340 | return af::SUCCESS; |
| 259 | } | 341 | } |
| 260 | 342 | ||
| 261 | -af::Status ClassifyIndirectLoadLayout(const LogicalTensorView &logical, IndirectLoadTensorLayout &layout) { | 343 | +af::Status ClassifyIndirectLoadLayout(const LogicalTensorView &logical, IndirectLoadTensorLayout &layout, |
| 344 | + bool allow_non_overlapping_zero_stride) { | ||
| 262 | GE_ASSERT_TRUE(IsValidLogicalTensorView(logical), "IndirectLoad input layout rank is invalid."); | 345 | GE_ASSERT_TRUE(IsValidLogicalTensorView(logical), "IndirectLoad input layout rank is invalid."); |
| 263 | static_cast<LogicalTensorView &>(layout) = logical; | 346 | static_cast<LogicalTensorView &>(layout) = logical; |
| 264 | layout.kind = IndirectLoadLayoutKind::kUnsupported; | 347 | layout.kind = IndirectLoadLayoutKind::kUnsupported; |
| @@ -289,6 +372,10 @@ af::Status ClassifyIndirectLoadLayout(const LogicalTensorView &logical, Indirect | |||
| 289 | physical_span = physical_span + (logical.sizes[dim] - af::sym::kSymbolOne) * logical.strides[dim]; | 372 | physical_span = physical_span + (logical.sizes[dim] - af::sym::kSymbolOne) * logical.strides[dim]; |
| 290 | } | 373 | } |
| 291 | if (has_zero_stride && has_physical_gap) { | 374 | if (has_zero_stride && has_physical_gap) { |
| 375 | + if (allow_non_overlapping_zero_stride) { | ||
| 376 | + layout.kind = IndirectLoadLayoutKind::kStrided; | ||
| 377 | + layout.physical_repeats = logical.sizes; | ||
| 378 | + } | ||
| 292 | return af::SUCCESS; | 379 | return af::SUCCESS; |
| 293 | } | 380 | } |
| 294 | layout.kind = has_zero_stride | 381 | layout.kind = has_zero_stride |
| @@ -31,6 +31,7 @@ enum class TemplateRole : int64_t { | |||
| 31 | kSimtInputBoundary, | 31 | kSimtInputBoundary, |
| 32 | kSimtDirectGmBoundary, | 32 | kSimtDirectGmBoundary, |
| 33 | kSimtInlineTransform, | 33 | kSimtInlineTransform, |
| 34 | + kSimtFanoutBranch, | ||
| 34 | kSimtOp, | 35 | kSimtOp, |
| 35 | kSkInputBoundary, | 36 | kSkInputBoundary, |
| 36 | kStridedUbPath, | 37 | kStridedUbPath, |
| @@ -100,7 +101,8 @@ af::Status SetTemplateLogicalView(const af::AscNodePtr &node, const TemplateLogi | |||
| 100 | af::Status GetTemplateLogicalView(const af::AscNodePtr &node, TemplateLogicalView &view); | 101 | af::Status GetTemplateLogicalView(const af::AscNodePtr &node, TemplateLogicalView &view); |
| 101 | af::Status SetImplementation(const af::AscNodePtr &node, Implementation implementation); | 102 | af::Status SetImplementation(const af::AscNodePtr &node, Implementation implementation); |
| 102 | af::Status GetImplementation(const af::AscNodePtr &node, Implementation &implementation); | 103 | af::Status GetImplementation(const af::AscNodePtr &node, Implementation &implementation); |
| 103 | -af::Status ClassifyIndirectLoadLayout(const LogicalTensorView &logical, IndirectLoadTensorLayout &layout); | 104 | +af::Status ClassifyIndirectLoadLayout(const LogicalTensorView &logical, IndirectLoadTensorLayout &layout, |
| 105 | + bool allow_non_overlapping_zero_stride = false); | ||
| 104 | af::Status ValidateIndirectLoadOutputLayout(const LogicalTensorView &output); | 106 | af::Status ValidateIndirectLoadOutputLayout(const LogicalTensorView &output); |
| 105 | bool ShouldApplyInputInnerVectorization(const af::AscNodePtr &node); | 107 | bool ShouldApplyInputInnerVectorization(const af::AscNodePtr &node); |
| 106 | af::AscNodePtr GetInputProducer(const af::AscNodePtr &node, size_t input_index); | 108 | af::AscNodePtr GetInputProducer(const af::AscNodePtr &node, size_t input_index); |
| @@ -550,6 +550,19 @@ Status ApplyIndirectLoadTemplateMerge(ascir::ImplGraph &graph, const af::AscNode | |||
| 550 | return af::SUCCESS; | 550 | return af::SUCCESS; |
| 551 | } | 551 | } |
| 552 | 552 | ||
| 553 | +Status SetOuterRepeatsToOne(const af::AscNodePtr &node, const std::vector<ascir::AxisId> &vectorized_axes) { | ||
| 554 | + if (node->attr.api.compute_type != af::ComputeType::kComputeReduce) { | ||
| 555 | + return af::SUCCESS; | ||
| 556 | + } | ||
| 557 | + GE_ASSERT_TRUE(!vectorized_axes.empty(), "Node[%s] has no IndirectLoad vector axes.", node->GetNamePtr()); | ||
| 558 | + for (const auto &output : node->outputs()) { | ||
| 559 | + GE_ASSERT_TRUE(output->attr.axis.size() == output->attr.repeats.size()); | ||
| 560 | + GE_ASSERT_TRUE(output->attr.axis.size() >= vectorized_axes.size()); | ||
| 561 | + std::fill_n(output->attr.repeats.begin(), output->attr.axis.size() - vectorized_axes.size(), af::sym::kSymbolOne); | ||
| 562 | + } | ||
| 563 | + return af::SUCCESS; | ||
| 564 | +} | ||
| 565 | + | ||
| 553 | Status SetInputInnerVectorizedView(const af::AscNodePtr &node, const ascir::Axis &input_inner_axis, | 566 | Status SetInputInnerVectorizedView(const af::AscNodePtr &node, const ascir::Axis &input_inner_axis, |
| 554 | af::AscTensorAttr &attr) { | 567 | af::AscTensorAttr &attr) { |
| 555 | GE_ASSERT_TRUE(attr.axis.size() == attr.repeats.size() && attr.axis.size() == attr.strides.size(), | 568 | GE_ASSERT_TRUE(attr.axis.size() == attr.repeats.size() && attr.axis.size() == attr.strides.size(), |
| @@ -607,40 +620,6 @@ Status ApplyInputInnerVectorizedAxis(ascir::ImplGraph &graph, const af::AscNodeP | |||
| 607 | return af::SUCCESS; | 620 | return af::SUCCESS; |
| 608 | } | 621 | } |
| 609 | 622 | ||
| 610 | -Status SetOuterRepeatsToOne(const af::AscNodePtr &node, const std::vector<ascir::AxisId> &vectorized_axes) { | ||
| 611 | - GE_ASSERT_TRUE(!vectorized_axes.empty(), "Node[%s] has no IndirectLoad vector axes.", node->GetNamePtr()); | ||
| 612 | - for (const auto &output : node->outputs()) { | ||
| 613 | - GE_ASSERT_TRUE(output->attr.axis.size() == output->attr.repeats.size()); | ||
| 614 | - GE_ASSERT_TRUE(output->attr.axis.size() >= vectorized_axes.size()); | ||
| 615 | - std::fill_n(output->attr.repeats.begin(), output->attr.axis.size() - vectorized_axes.size(), af::sym::kSymbolOne); | ||
| 616 | - } | ||
| 617 | - return af::SUCCESS; | ||
| 618 | -} | ||
| 619 | - | ||
| 620 | -Status SetReduceInputVectorizedView(const af::AscNodePtr &reduce, const af::AscNodePtr &input_producer, | ||
| 621 | - const std::vector<ascir::AxisId> &vectorized_axes) { | ||
| 622 | - GE_ASSERT_TRUE(!vectorized_axes.empty(), "Reduce node[%s] has no IndirectLoad vector axes.", reduce->GetNamePtr()); | ||
| 623 | - for (size_t i = 0UL; i < reduce->inputs.Size(); ++i) { | ||
| 624 | - if (ascgen_utils::indirect_load::GetInputProducer(reduce, i) != input_producer) { | ||
| 625 | - continue; | ||
| 626 | - } | ||
| 627 | - auto &input = reduce->inputs[i].attr; | ||
| 628 | - GE_ASSERT_TRUE(input.axis.size() == input.strides.size()); | ||
| 629 | - input.vectorized_axis = vectorized_axes; | ||
| 630 | - input.vectorized_strides.clear(); | ||
| 631 | - input.vectorized_strides.reserve(vectorized_axes.size()); | ||
| 632 | - for (ascir::AxisId axis : vectorized_axes) { | ||
| 633 | - const auto iter = std::find(input.axis.begin(), input.axis.end(), axis); | ||
| 634 | - GE_ASSERT_TRUE(iter != input.axis.end(), "Reduce node[%s] has no IndirectLoad vector axis[%ld].", | ||
| 635 | - reduce->GetNamePtr(), axis); | ||
| 636 | - input.vectorized_strides.emplace_back( | ||
| 637 | - input.strides[static_cast<size_t>(std::distance(input.axis.begin(), iter))]); | ||
| 638 | - } | ||
| 639 | - return af::SUCCESS; | ||
| 640 | - } | ||
| 641 | - GELOGE(af::FAILED, "Reduce node[%s] is not connected to IndirectLoad output path.", reduce->GetNamePtr()); | ||
| 642 | - return af::FAILED; | ||
| 643 | -} | ||
| 644 | } // namespace | 623 | } // namespace |
| 645 | 624 | ||
| 646 | Status Scheduler::InitIndirectLoadScheduleCase() { | 625 | Status Scheduler::InitIndirectLoadScheduleCase() { |
| @@ -659,11 +638,6 @@ Status Scheduler::InitIndirectLoadScheduleCase() { | |||
| 659 | } | 638 | } |
| 660 | GE_ASSERT_NOTNULL(tiling_case_.ub_tiling_y.first); | 639 | GE_ASSERT_NOTNULL(tiling_case_.ub_tiling_y.first); |
| 661 | GE_ASSERT_SUCCESS(ascgen_utils::indirect_load::GetTemplateAxes(indirect_load, indirect_load_info_.axes)); | 640 | GE_ASSERT_SUCCESS(ascgen_utils::indirect_load::GetTemplateAxes(indirect_load, indirect_load_info_.axes)); |
| 662 | - indirect_load_info_.reduce = ascgen_utils::indirect_load::GetPostReduceConsumer(indirect_load); | ||
| 663 | - if (indirect_load_info_.reduce != nullptr) { | ||
| 664 | - indirect_load_info_.reduce_input = ascgen_utils::indirect_load::GetPostReduceInputProducer(indirect_load); | ||
| 665 | - GE_ASSERT_NOTNULL(indirect_load_info_.reduce_input, "IndirectLoad post Reduce input producer is missing."); | ||
| 666 | - } | ||
| 667 | indirect_load_info_.active = true; | 641 | indirect_load_info_.active = true; |
| 668 | return af::SUCCESS; | 642 | return af::SUCCESS; |
| 669 | } | 643 | } |
| @@ -694,11 +668,7 @@ Status Scheduler::ApplyIndirectLoadNodeAxes(const af::AscNodePtr &node, bool &sk | |||
| 694 | return af::SUCCESS; | 668 | return af::SUCCESS; |
| 695 | } | 669 | } |
| 696 | GE_ASSERT_SUCCESS(ApplyIndirectLoadTemplateMerge(graph_, node, indirect_load_info_.axes.inner_axis, false)); | 670 | GE_ASSERT_SUCCESS(ApplyIndirectLoadTemplateMerge(graph_, node, indirect_load_info_.axes.inner_axis, false)); |
| 697 | - if (node == indirect_load_info_.reduce) { | 671 | + GE_ASSERT_SUCCESS(SetOuterRepeatsToOne(node, indirect_load_info_.axes.vectorized_axes)); |
| 698 | - GE_ASSERT_SUCCESS(SetOuterRepeatsToOne(node, indirect_load_info_.axes.vectorized_axes)); | ||
| 699 | - GE_ASSERT_SUCCESS( | ||
| 700 | - SetReduceInputVectorizedView(node, indirect_load_info_.reduce_input, indirect_load_info_.axes.vectorized_axes)); | ||
| 701 | - } | ||
| 702 | return af::SUCCESS; | 672 | return af::SUCCESS; |
| 703 | } | 673 | } |
| 704 | 674 | ||
| @@ -43,8 +43,6 @@ struct TilingCase { | |||
| 43 | 43 | ||
| 44 | struct IndirectLoadInfo { | 44 | struct IndirectLoadInfo { |
| 45 | bool active = false; | 45 | bool active = false; |
| 46 | - af::AscNodePtr reduce; | ||
| 47 | - af::AscNodePtr reduce_input; | ||
| 48 | std::vector<af::AscNodePtr> aligned_strided_path; | 46 | std::vector<af::AscNodePtr> aligned_strided_path; |
| 49 | ascgen_utils::indirect_load::TemplateAxes axes; | 47 | ascgen_utils::indirect_load::TemplateAxes axes; |
| 50 | }; | 48 | }; |
| @@ -197,8 +197,8 @@ af::Status BuildBroadcastLogicalView(const af::AscTensorAttr &logical_attr, cons | |||
| 197 | return af::SUCCESS; | 197 | return af::SUCCESS; |
| 198 | } | 198 | } |
| 199 | 199 | ||
| 200 | -af::Status ApplyZeroStrideCompactView(const NodePath &path, | 200 | +af::Status ApplyPhysicalView(const NodePath &path, |
| 201 | - const ascgen_utils::indirect_load::IndirectLoadTensorLayout &layout) { | 201 | + const ascgen_utils::indirect_load::IndirectLoadTensorLayout &layout) { |
| 202 | GE_ASSERT_TRUE( | 202 | GE_ASSERT_TRUE( |
| 203 | layout.axis_ids.size() == layout.physical_repeats.size() && layout.axis_ids.size() == layout.strides.size(), | 203 | layout.axis_ids.size() == layout.physical_repeats.size() && layout.axis_ids.size() == layout.strides.size(), |
| 204 | "IndirectLoad physical execution view rank mismatch."); | 204 | "IndirectLoad physical execution view rank mismatch."); |
| @@ -255,7 +255,18 @@ af::Status ApplyIndirectLoadPathLayout(const NodePath &path, | |||
| 255 | const ascgen_utils::indirect_load::IndirectLoadTensorLayout &layout, | 255 | const ascgen_utils::indirect_load::IndirectLoadTensorLayout &layout, |
| 256 | bool needs_alignment) { | 256 | bool needs_alignment) { |
| 257 | if (layout.kind == ascgen_utils::indirect_load::IndirectLoadLayoutKind::kZeroStrideCompact) { | 257 | if (layout.kind == ascgen_utils::indirect_load::IndirectLoadLayoutKind::kZeroStrideCompact) { |
| 258 | - return ApplyZeroStrideCompactView(path, layout); | 258 | + return ApplyPhysicalView(path, layout); |
| 259 | + } | ||
| 260 | + // Use the rewritten path as the source of truth; Broadcast may have been folded before this point. | ||
| 261 | + const auto broadcast = | ||
| 262 | + std::find_if(path.begin(), path.end(), [](const af::AscNodePtr &node) { return IsBroadcastNode(node); }); | ||
| 263 | + const bool has_dynamic_shape = std::any_of(layout.sizes.begin(), layout.sizes.end(), | ||
| 264 | + [](const af::Expression &size) { return !size.IsConstExpr(); }); | ||
| 265 | + if (layout.kind == ascgen_utils::indirect_load::IndirectLoadLayoutKind::kStrided && has_dynamic_shape && | ||
| 266 | + broadcast != path.end()) { | ||
| 267 | + // Keep the upstream producer view intact; only the rewritten producer-to-Broadcast segment is normalized. | ||
| 268 | + const NodePath broadcast_path(path.begin(), broadcast + 1); | ||
| 269 | + GE_ASSERT_SUCCESS(ApplyPhysicalView(broadcast_path, layout)); | ||
| 259 | } | 270 | } |
| 260 | if (needs_alignment) { | 271 | if (needs_alignment) { |
| 261 | return AnnotateStridedUbPath(path); | 272 | return AnnotateStridedUbPath(path); |
| @@ -779,6 +790,13 @@ af::Status NormalizeAxesForTemplate(af::AscGraph &graph, const af::AscNodePtr &i | |||
| 779 | GE_ASSERT_NOTNULL(inner, "IndirectLoad inner axis %ld is not found.", inner_axis); | 790 | GE_ASSERT_NOTNULL(inner, "IndirectLoad inner axis %ld is not found.", inner_axis); |
| 780 | vectorized_axes = | 791 | vectorized_axes = |
| 781 | inner->type == ascir::Axis::Type::kAxisTypeMerged ? inner->from : std::vector<af::AxisId>{inner_axis}; | 792 | inner->type == ascir::Axis::Type::kAxisTypeMerged ? inner->from : std::vector<af::AxisId>{inner_axis}; |
| 793 | + } else { | ||
| 794 | + // A SIMT candidate without post Reduce uses the complete output view as | ||
| 795 | + // the outer view. The IndirectLoad/its fused direct-GM path intentionally | ||
| 796 | + // keeps an empty vectorized view, while ordinary fan-out branches still | ||
| 797 | + // need the tile-inner axis for alignment and vector-function partitioning. | ||
| 798 | + GE_ASSERT_TRUE(tile_inner_axis != af::kIdNone, "IndirectLoad tile inner axis is missing."); | ||
| 799 | + vectorized_axes.emplace_back(tile_inner_axis); | ||
| 782 | } | 800 | } |
| 783 | ascgen_utils::indirect_load::TemplateAxes axes; | 801 | ascgen_utils::indirect_load::TemplateAxes axes; |
| 784 | axes.outer_axis = outer_axis; | 802 | axes.outer_axis = outer_axis; |
| @@ -793,6 +811,28 @@ af::Status NormalizeAxesForTemplate(af::AscGraph &graph, const af::AscNodePtr &i | |||
| 793 | return af::SUCCESS; | 811 | return af::SUCCESS; |
| 794 | } | 812 | } |
| 795 | 813 | ||
| 814 | +af::Status SeedPostReduceInputVectorizedView(const af::AscNodePtr &indirect_load, const af::AscNodePtr &reduce) { | ||
| 815 | + if (reduce == nullptr) { | ||
| 816 | + return af::SUCCESS; | ||
| 817 | + } | ||
| 818 | + ascgen_utils::indirect_load::TemplateAxes axes; | ||
| 819 | + GE_ASSERT_SUCCESS(ascgen_utils::indirect_load::GetTemplateAxes(indirect_load, axes)); | ||
| 820 | + GE_ASSERT_TRUE(!axes.vectorized_axes.empty(), "IndirectLoad post Reduce vectorized axes are empty."); | ||
| 821 | + GE_ASSERT_TRUE(reduce->inputs.Size() == 1UL, "IndirectLoad post Reduce must have one input."); | ||
| 822 | + auto &input = reduce->inputs[0].attr; | ||
| 823 | + GE_ASSERT_TRUE(input.axis.size() == input.strides.size(), "IndirectLoad post Reduce input view is invalid."); | ||
| 824 | + input.vectorized_axis = axes.vectorized_axes; | ||
| 825 | + input.vectorized_strides.clear(); | ||
| 826 | + input.vectorized_strides.reserve(input.vectorized_axis.size()); | ||
| 827 | + for (const auto axis : input.vectorized_axis) { | ||
| 828 | + const auto axis_it = std::find(input.axis.begin(), input.axis.end(), axis); | ||
| 829 | + GE_ASSERT_TRUE(axis_it != input.axis.end(), "IndirectLoad post Reduce input axis[%ld] is missing.", axis); | ||
| 830 | + input.vectorized_strides.emplace_back( | ||
| 831 | + input.strides[static_cast<size_t>(std::distance(input.axis.begin(), axis_it))]); | ||
| 832 | + } | ||
| 833 | + return af::SUCCESS; | ||
| 834 | +} | ||
| 835 | + | ||
| 796 | af::Status RestoreSkTemplateAxes(std::vector<ascir::ImplGraph> &grouped_graphs) { | 836 | af::Status RestoreSkTemplateAxes(std::vector<ascir::ImplGraph> &grouped_graphs) { |
| 797 | for (auto &graph : grouped_graphs) { | 837 | for (auto &graph : grouped_graphs) { |
| 798 | af::AscNodePtr indirect_load; | 838 | af::AscNodePtr indirect_load; |
| @@ -869,27 +909,104 @@ af::Status NormalizeSimtAxesForTemplate(af::AscGraph &graph, const af::AscNodePt | |||
| 869 | return NormalizeAxesForTemplate(graph, indirect_load, boundary, af::kIdNone, af::kIdNone); | 909 | return NormalizeAxesForTemplate(graph, indirect_load, boundary, af::kIdNone, af::kIdNone); |
| 870 | } | 910 | } |
| 871 | 911 | ||
| 912 | +// SIMT direct-GM nodes deliberately preserve their vectorized view through the | ||
| 913 | +// scheduler. User graphs, however, do not necessarily initialize that view; | ||
| 914 | +// fill it from the node's physical strides before VF partitioning consumes it. | ||
| 915 | +af::Status CompletePreservedVectorizedViews(const af::AscGraph &graph, const af::AscNodePtr &indirect_load) { | ||
| 916 | + for (const auto &node : graph.GetAllNodes()) { | ||
| 917 | + if (node == nullptr || node == indirect_load || | ||
| 918 | + !ascgen_utils::indirect_load::GetTemplateBehavior(node).preserves_vectorized_axis) { | ||
| 919 | + continue; | ||
| 920 | + } | ||
| 921 | + for (const auto &output : node->outputs()) { | ||
| 922 | + if (output == nullptr || !output->attr.vectorized_axis.empty()) { | ||
| 923 | + continue; | ||
| 924 | + } | ||
| 925 | + const auto &axes = output->attr.axis; | ||
| 926 | + const auto &strides = output->attr.strides; | ||
| 927 | + GE_ASSERT_TRUE(!axes.empty() && axes.size() == strides.size(), | ||
| 928 | + "IndirectLoad preserved node[%s] has invalid tensor view.", node->GetNamePtr()); | ||
| 929 | + size_t vectorized_index = axes.size() - 1UL; | ||
| 930 | + for (size_t index = axes.size(); index > 0UL; --index) { | ||
| 931 | + const auto stride = strides[index - 1UL]; | ||
| 932 | + if (af::SymbolicUtils::StaticCheckEq(stride, af::sym::kSymbolOne) == af::TriBool::kTrue) { | ||
| 933 | + vectorized_index = index - 1UL; | ||
| 934 | + break; | ||
| 935 | + } | ||
| 936 | + if (af::SymbolicUtils::StaticCheckEq(stride, af::sym::kSymbolZero) != af::TriBool::kTrue) { | ||
| 937 | + vectorized_index = index - 1UL; | ||
| 938 | + } | ||
| 939 | + } | ||
| 940 | + output->attr.vectorized_axis = {axes[vectorized_index]}; | ||
| 941 | + output->attr.vectorized_strides = {strides[vectorized_index]}; | ||
| 942 | + GELOGD("[IndirectLoad] Seed preserved vectorized view for node[%s], axis[%ld].", node->GetNamePtr(), | ||
| 943 | + axes[vectorized_index]); | ||
| 944 | + } | ||
| 945 | + } | ||
| 946 | + return af::SUCCESS; | ||
| 947 | +} | ||
| 948 | + | ||
| 872 | bool CanEmitSimtScalar(const af::AscNodePtr &node) { | 949 | bool CanEmitSimtScalar(const af::AscNodePtr &node) { |
| 873 | const auto impl = ascgen_utils::GetAscIrCodegenImpl(node->GetType()); | 950 | const auto impl = ascgen_utils::GetAscIrCodegenImpl(node->GetType()); |
| 874 | const auto *v2_impl = impl == nullptr ? nullptr : dynamic_cast<af::ascir::AscIrCodegenV2 *>(impl.get()); | 951 | const auto *v2_impl = impl == nullptr ? nullptr : dynamic_cast<af::ascir::AscIrCodegenV2 *>(impl.get()); |
| 875 | return v2_impl != nullptr && v2_impl->IsSimtScalarSupported(*node); | 952 | return v2_impl != nullptr && v2_impl->IsSimtScalarSupported(*node); |
| 876 | } | 953 | } |
| 877 | 954 | ||
| 878 | -// 一次下游行走同时拿输出链上的 Store 与 Reduce;多于一个 Reduce 时没有任何模板能支持,直接断言报错。 | 955 | +af::Status ValidateReduceOutput(const af::AscNodePtr &reduce) { |
| 956 | + const auto reduce_outputs = reduce->GetOutDataNodes(); | ||
| 957 | + GE_ASSERT_TRUE(reduce_outputs.size() == 1UL, "[IndirectLoad] Reduce[%s] must have one Store output, got %zu.", | ||
| 958 | + reduce->GetNamePtr(), reduce_outputs.size()); | ||
| 959 | + const auto successor = std::dynamic_pointer_cast<af::AscNode>(*reduce_outputs.begin()); | ||
| 960 | + GE_ASSERT_NOTNULL(successor, "IndirectLoad Reduce output node is invalid."); | ||
| 961 | + if (ScheduleUtils::IsStore(successor)) { | ||
| 962 | + return af::SUCCESS; | ||
| 963 | + } | ||
| 964 | + GE_ASSERT_TRUE(af::ops::IsOps<af::ascir_op::Cast>(successor), | ||
| 965 | + "[IndirectLoad] Reduce[%s] output must be Store or Cast, got node[%s] type[%s].", reduce->GetNamePtr(), | ||
| 966 | + successor->GetNamePtr(), successor->GetTypePtr()); | ||
| 967 | + const auto cast_outputs = successor->GetOutDataNodes(); | ||
| 968 | + GE_ASSERT_TRUE(cast_outputs.size() == 1UL, "[IndirectLoad] Reduce[%s] Cast must have one Store output, got %zu.", | ||
| 969 | + reduce->GetNamePtr(), cast_outputs.size()); | ||
| 970 | + const auto cast_successor = std::dynamic_pointer_cast<af::AscNode>(*cast_outputs.begin()); | ||
| 971 | + GE_ASSERT_NOTNULL(cast_successor, "IndirectLoad Reduce Cast output node is invalid."); | ||
| 972 | + GE_ASSERT_TRUE(ScheduleUtils::IsStore(cast_successor), | ||
| 973 | + "[IndirectLoad] Reduce[%s] Cast output must be Store, got node[%s] type[%s].", reduce->GetNamePtr(), | ||
| 974 | + cast_successor->GetNamePtr(), cast_successor->GetTypePtr()); | ||
| 975 | + return af::SUCCESS; | ||
| 976 | +} | ||
| 977 | + | ||
| 978 | +// Traverse all output branches once, collecting Store/Reduce boundaries and validating each Reduce successor. | ||
| 879 | af::Status CollectOutputBoundaries(const af::AscNodePtr &indirect_load, RewrittenGraphAnalysis &analysis) { | 979 | af::Status CollectOutputBoundaries(const af::AscNodePtr &indirect_load, RewrittenGraphAnalysis &analysis) { |
| 880 | - for (af::AscNodePtr current = ascgen_utils::indirect_load::GetOnlyOutputConsumer(indirect_load); current != nullptr; | 980 | + NodePath pending; |
| 881 | - current = ascgen_utils::indirect_load::GetOnlyOutputConsumer(current)) { | 981 | + NodeSet visited{indirect_load.get()}; |
| 882 | - if (analysis.output_store == nullptr && af::ops::IsOps<af::ascir_op::Store>(current)) { | 982 | + for (const auto &out_node : indirect_load->GetOutDataNodes()) { |
| 883 | - analysis.output_store = current; | 983 | + const auto out_asc_node = std::dynamic_pointer_cast<af::AscNode>(out_node); |
| 984 | + GE_ASSERT_NOTNULL(out_asc_node, "IndirectLoad output successor is invalid."); | ||
| 985 | + if (visited.emplace(out_asc_node.get()).second) { | ||
| 986 | + pending.emplace_back(out_asc_node); | ||
| 884 | } | 987 | } |
| 885 | - if (!ScheduleUtils::IsReduce(current)) { | 988 | + } |
| 989 | + for (size_t cursor = 0UL; cursor < pending.size(); ++cursor) { | ||
| 990 | + const auto ¤t = pending[cursor]; | ||
| 991 | + if (ScheduleUtils::IsStore(current)) { | ||
| 992 | + if (analysis.output_store == nullptr) { | ||
| 993 | + analysis.output_store = current; | ||
| 994 | + } | ||
| 886 | continue; | 995 | continue; |
| 887 | } | 996 | } |
| 888 | - GE_ASSERT_TRUE( | 997 | + if (ScheduleUtils::IsReduce(current)) { |
| 889 | - analysis.post_reduce == nullptr, | 998 | + GE_ASSERT_SUCCESS(ValidateReduceOutput(current)); |
| 890 | - "[IndirectLoad] IndirectLoad node[%s] post chain contains more than one Reduce: node[%s] and node[%s].", | 999 | + GE_ASSERT_TRUE(analysis.post_reduce == nullptr); |
| 891 | - indirect_load->GetNamePtr(), analysis.post_reduce->GetNamePtr(), current->GetNamePtr()); | 1000 | + analysis.post_reduce = current; |
| 892 | - analysis.post_reduce = current; | 1001 | + continue; |
| 1002 | + } | ||
| 1003 | + for (const auto &out_node : current->GetOutDataNodes()) { | ||
| 1004 | + const auto out_asc_node = std::dynamic_pointer_cast<af::AscNode>(out_node); | ||
| 1005 | + GE_ASSERT_NOTNULL(out_asc_node, "IndirectLoad output successor is invalid."); | ||
| 1006 | + if (visited.emplace(out_asc_node.get()).second) { | ||
| 1007 | + pending.emplace_back(out_asc_node); | ||
| 1008 | + } | ||
| 1009 | + } | ||
| 893 | } | 1010 | } |
| 894 | return af::SUCCESS; | 1011 | return af::SUCCESS; |
| 895 | } | 1012 | } |
| @@ -1131,7 +1248,9 @@ af::Status AnalyzeInputPath(const af::AscNodePtr &indirect_load, size_t input_id | |||
| 1131 | GE_ASSERT_SUCCESS(BuildBroadcastLogicalView(logical_attr, physical_attr, template_id, view)); | 1248 | GE_ASSERT_SUCCESS(BuildBroadcastLogicalView(logical_attr, physical_attr, template_id, view)); |
| 1132 | } | 1249 | } |
| 1133 | 1250 | ||
| 1134 | - GE_ASSERT_SUCCESS(ascgen_utils::indirect_load::ClassifyIndirectLoadLayout(view, plan.layout)); | 1251 | + const bool allow_simt_strided_broadcast = template_id == ascir::TemplateId::kIndirectLoadSimt && has_broadcast; |
| 1252 | + GE_ASSERT_SUCCESS( | ||
| 1253 | + ascgen_utils::indirect_load::ClassifyIndirectLoadLayout(view, plan.layout, allow_simt_strided_broadcast)); | ||
| 1135 | if (plan.layout.kind == ascgen_utils::indirect_load::IndirectLoadLayoutKind::kUnsupported) { | 1254 | if (plan.layout.kind == ascgen_utils::indirect_load::IndirectLoadLayoutKind::kUnsupported) { |
| 1136 | GELOGI("[IndirectLoad] Reject candidate[%d]: input path layout%s is unsupported.", | 1255 | GELOGI("[IndirectLoad] Reject candidate[%d]: input path layout%s is unsupported.", |
| 1137 | static_cast<int32_t>(template_id), has_broadcast ? " with Broadcast source view" : ""); | 1256 | static_cast<int32_t>(template_id), has_broadcast ? " with Broadcast source view" : ""); |
| @@ -1309,10 +1428,9 @@ af::Status ValidateSimtTemplateRegion(const RewrittenGraphAnalysis &analysis, bo | |||
| 1309 | return af::SUCCESS; | 1428 | return af::SUCCESS; |
| 1310 | } | 1429 | } |
| 1311 | for (const af::AscNodePtr &node : analysis.index_region) { | 1430 | for (const af::AscNodePtr &node : analysis.index_region) { |
| 1312 | - if (af::ops::IsOps<af::ascir_op::ScalarData>(node) || af::ops::IsOps<af::ascir_op::Scalar>(node)) { | 1431 | + // Compile-time Scalar values are emitted as local constants by the SIMT evaluator. |
| 1313 | - return af::SUCCESS; | 1432 | + // ScalarData is runtime input and still requires an explicit context/GM binding. |
| 1314 | - } | 1433 | + if (af::ops::IsOps<af::ascir_op::ScalarData>(node) || HasControlEdge(node)) { |
| 1315 | - if (HasControlEdge(node) || af::ops::IsOps<af::ascir_op::VectorFunc>(node)) { | ||
| 1316 | return af::SUCCESS; | 1434 | return af::SUCCESS; |
| 1317 | } | 1435 | } |
| 1318 | const bool is_gm_boundary = af::ops::IsOps<af::ascir_op::Load>(node) || af::ops::IsOps<af::ascir_op::Store>(node); | 1436 | const bool is_gm_boundary = af::ops::IsOps<af::ascir_op::Load>(node) || af::ops::IsOps<af::ascir_op::Store>(node); |
| @@ -1340,6 +1458,48 @@ af::Status AnnotateSimtTemplateRoles(const af::AscNodePtr &indirect_load, const | |||
| 1340 | return af::SUCCESS; | 1458 | return af::SUCCESS; |
| 1341 | } | 1459 | } |
| 1342 | 1460 | ||
| 1461 | +af::Status AnnotateSimtFanoutBranches(const af::AscNodePtr &indirect_load, const RewrittenGraphAnalysis &analysis) { | ||
| 1462 | + const af::AscNodePtr selected_root = analysis.post_reduce == nullptr | ||
| 1463 | + ? analysis.output_store | ||
| 1464 | + : ascgen_utils::indirect_load::GetInputProducer(analysis.post_reduce, 0UL); | ||
| 1465 | + if (selected_root == nullptr) { | ||
| 1466 | + return af::SUCCESS; | ||
| 1467 | + } | ||
| 1468 | + NodeSet selected; | ||
| 1469 | + const bool selected_ok = CollectSimtBackwardRegion({selected_root}, indirect_load, selected); | ||
| 1470 | + if (!selected_ok) { | ||
| 1471 | + GELOGW("[IndirectLoad] Cannot identify SIMT main output chain from root[%s]; keep existing roles.", | ||
| 1472 | + selected_root->GetNamePtr()); | ||
| 1473 | + return af::SUCCESS; | ||
| 1474 | + } | ||
| 1475 | + | ||
| 1476 | + NodePath pending; | ||
| 1477 | + NodeSet visited{indirect_load.get()}; | ||
| 1478 | + for (const auto &out_node : indirect_load->GetOutDataNodes()) { | ||
| 1479 | + const auto consumer = std::dynamic_pointer_cast<af::AscNode>(out_node); | ||
| 1480 | + if (consumer != nullptr && visited.emplace(consumer.get()).second) { | ||
| 1481 | + pending.emplace_back(consumer); | ||
| 1482 | + } | ||
| 1483 | + } | ||
| 1484 | + for (size_t cursor = 0UL; cursor < pending.size(); ++cursor) { | ||
| 1485 | + const auto &node = pending[cursor]; | ||
| 1486 | + if (selected.count(node.get()) == 0UL && !IsInputRegionBoundary(node) && !ScheduleUtils::IsReduce(node)) { | ||
| 1487 | + GE_ASSERT_SUCCESS(ascgen_utils::indirect_load::SetTemplateRole( | ||
| 1488 | + node, ascgen_utils::indirect_load::TemplateRole::kSimtFanoutBranch)); | ||
| 1489 | + } | ||
| 1490 | + if (ScheduleUtils::IsStore(node) || ScheduleUtils::IsReduce(node)) { | ||
| 1491 | + continue; | ||
| 1492 | + } | ||
| 1493 | + for (const auto &out_node : node->GetOutDataNodes()) { | ||
| 1494 | + const auto consumer = std::dynamic_pointer_cast<af::AscNode>(out_node); | ||
| 1495 | + if (consumer != nullptr && visited.emplace(consumer.get()).second) { | ||
| 1496 | + pending.emplace_back(consumer); | ||
| 1497 | + } | ||
| 1498 | + } | ||
| 1499 | + } | ||
| 1500 | + return af::SUCCESS; | ||
| 1501 | +} | ||
| 1502 | + | ||
| 1343 | af::Status ValidateTemplate(const af::AscNodePtr &indirect_load, ascir::TemplateId template_id, | 1503 | af::Status ValidateTemplate(const af::AscNodePtr &indirect_load, ascir::TemplateId template_id, |
| 1344 | const RewrittenGraphAnalysis &analysis, size_t &boundary, bool &is_candidate_legal) { | 1504 | const RewrittenGraphAnalysis &analysis, size_t &boundary, bool &is_candidate_legal) { |
| 1345 | if (template_id == ascir::TemplateId::kIndirectLoadSimd) { | 1505 | if (template_id == ascir::TemplateId::kIndirectLoadSimd) { |
| @@ -1367,10 +1527,11 @@ af::Status NormalizeTemplateAxes(af::AscGraph &graph, const af::AscNodePtr &indi | |||
| 1367 | ascir::TemplateId template_id, const RewrittenGraphAnalysis &analysis, | 1527 | ascir::TemplateId template_id, const RewrittenGraphAnalysis &analysis, |
| 1368 | size_t boundary) { | 1528 | size_t boundary) { |
| 1369 | if (template_id == ascir::TemplateId::kIndirectLoadSimd) { | 1529 | if (template_id == ascir::TemplateId::kIndirectLoadSimd) { |
| 1370 | - return NormalizeSimdAxesForTemplate(graph, indirect_load, analysis); | 1530 | + GE_ASSERT_SUCCESS(NormalizeSimdAxesForTemplate(graph, indirect_load, analysis)); |
| 1371 | } else { | 1531 | } else { |
| 1372 | - return NormalizeSimtAxesForTemplate(graph, indirect_load, boundary); | 1532 | + GE_ASSERT_SUCCESS(NormalizeSimtAxesForTemplate(graph, indirect_load, boundary)); |
| 1373 | } | 1533 | } |
| 1534 | + return SeedPostReduceInputVectorizedView(indirect_load, analysis.post_reduce); | ||
| 1374 | } | 1535 | } |
| 1375 | 1536 | ||
| 1376 | af::Status FinalizeTemplate(const af::AscNodePtr &indirect_load, ascir::TemplateId template_id) { | 1537 | af::Status FinalizeTemplate(const af::AscNodePtr &indirect_load, ascir::TemplateId template_id) { |
| @@ -1434,7 +1595,11 @@ af::Status ApplyGraphPass(af::AscGraph &graph, const af::AscNodePtr &indirect_lo | |||
| 1434 | return af::SUCCESS; | 1595 | return af::SUCCESS; |
| 1435 | } | 1596 | } |
| 1436 | GE_ASSERT_SUCCESS(AnnotateTemplate(indirect_load, template_id, analysis)); | 1597 | GE_ASSERT_SUCCESS(AnnotateTemplate(indirect_load, template_id, analysis)); |
| 1598 | + if (template_id == ascir::TemplateId::kIndirectLoadSimt) { | ||
| 1599 | + GE_ASSERT_SUCCESS(AnnotateSimtFanoutBranches(indirect_load, analysis)); | ||
| 1600 | + } | ||
| 1437 | GE_ASSERT_SUCCESS(NormalizeTemplateAxes(graph, indirect_load, template_id, analysis, boundary)); | 1601 | GE_ASSERT_SUCCESS(NormalizeTemplateAxes(graph, indirect_load, template_id, analysis, boundary)); |
| 1602 | + GE_ASSERT_SUCCESS(CompletePreservedVectorizedViews(graph, indirect_load)); | ||
| 1438 | return FinalizeTemplate(indirect_load, template_id); | 1603 | return FinalizeTemplate(indirect_load, template_id); |
| 1439 | } | 1604 | } |
| 1440 | 1605 | ||
| @@ -296,8 +296,8 @@ Status ReducePartitionCaseGenerator::GeneratorTask(ascir::HintGraph &optimize_gr | |||
| 296 | const OptimizerOptions &options) { | 296 | const OptimizerOptions &options) { |
| 297 | (void)options; | 297 | (void)options; |
| 298 | const af::AscNodePtr indirect_load = ascgen_utils::indirect_load::FindIndirectLoadNode(optimize_graph); | 298 | const af::AscNodePtr indirect_load = ascgen_utils::indirect_load::FindIndirectLoadNode(optimize_graph); |
| 299 | - if (indirect_load != nullptr && ascgen_utils::indirect_load::GetPostReduceConsumer(indirect_load) != nullptr) { | 299 | + if (indirect_load != nullptr) { |
| 300 | - GELOGI("Graph %s has indirect load with post reduce consumer, skip reduce task generation.", | 300 | + GELOGI("Graph %s has indirect load, skip reduce task generation; handled by IndirectLoad tasks.", |
| 301 | optimize_graph.GetName().c_str()); | 301 | optimize_graph.GetName().c_str()); |
| 302 | return ge::GRAPH_SUCCESS; | 302 | return ge::GRAPH_SUCCESS; |
| 303 | } | 303 | } |
| @@ -20,6 +20,7 @@ | |||
| 20 | 20 | ||
| 21 | 21 | ||
| 22 | 22 | ||
| 23 | + | ||
| 23 | 24 | ||
| 24 | 25 | ||
| 25 | using namespace af::ops; | 26 | using namespace af::ops; |
| @@ -29,6 +30,14 @@ using namespace ascgen_utils::indirect_load; | |||
| 29 | 30 | ||
| 30 | namespace { | 31 | namespace { |
| 31 | 32 | ||
| 33 | +TEST(IndirectLoadCodegenImplTest, LoadsApiHeadersInDeclarationOrder) { | ||
| 34 | + const auto headers = af::ascir::IndirectLoadAscIrCodegenImplV2().LoadApiHeaderFiles(false); | ||
| 35 | + const std::vector<std::string> expected = {"datacopy_reg_base.h", "indirect_load_simd_policy_reg_base.h", | ||
| 36 | + "indirect_load_simd_reg_base.h", "indirect_load_sk_reg_base.h", | ||
| 37 | + "indirect_load_simt_reg_base.h"}; | ||
| 38 | + EXPECT_EQ(headers, expected); | ||
| 39 | +} | ||
| 40 | + | ||
| 32 | struct ILTestGraph { | 41 | struct ILTestGraph { |
| 33 | af::AscGraph graph; | 42 | af::AscGraph graph; |
| 34 | af::Expression s0; | 43 | af::Expression s0; |
| @@ -871,15 +880,17 @@ TEST(IndirectLoadApiCallTest, GenerateSimdProducesIndirectLoadSimdCall) { | |||
| 871 | std::string result; | 880 | std::string result; |
| 872 | GenerateSimdCall(ge::DT_FLOAT16, result); | 881 | GenerateSimdCall(ge::DT_FLOAT16, result); |
| 873 | EXPECT_NE(result.find("// IndirectLoad SIMD"), std::string::npos); | 882 | EXPECT_NE(result.find("// IndirectLoad SIMD"), std::string::npos); |
| 874 | - EXPECT_NE(result.find("IndirectLoadSimd<half, int32_t, 2, 1>"), std::string::npos); | 883 | + EXPECT_NE(result.find("IndirectLoadSimdStrided<half, int32_t, 2, 1>"), std::string::npos); |
| 875 | // Symbolic sizes are dynamic shapes, which use the strided path and require a temporary UB buffer. | 884 | // Symbolic sizes are dynamic shapes, which use the strided path and require a temporary UB buffer. |
| 876 | EXPECT_NE(result.find("tmp_buf_0"), std::string::npos); | 885 | EXPECT_NE(result.find("tmp_buf_0"), std::string::npos); |
| 886 | + EXPECT_NE(result.find("indirect_load_simd_params{static_cast<uint32_t>("), std::string::npos); | ||
| 887 | + EXPECT_NE(result.find("), static_cast<uint32_t>("), std::string::npos); | ||
| 877 | } | 888 | } |
| 878 | 889 | ||
| 879 | TEST(IndirectLoadApiCallTest, GenerateSimdUint32InputProducesTypedCall) { | 890 | TEST(IndirectLoadApiCallTest, GenerateSimdUint32InputProducesTypedCall) { |
| 880 | std::string result; | 891 | std::string result; |
| 881 | GenerateSimdCall(ge::DT_UINT32, result); | 892 | GenerateSimdCall(ge::DT_UINT32, result); |
| 882 | - EXPECT_NE(result.find("IndirectLoadSimd<uint32_t, int32_t, 2, 1>"), std::string::npos); | 893 | + EXPECT_NE(result.find("IndirectLoadSimdStrided<uint32_t, int32_t, 2, 1>"), std::string::npos); |
| 883 | } | 894 | } |
| 884 | 895 | ||
| 885 | // ==================== SIMT Init + GenerateFuncDefinition ==================== | 896 | // ==================== SIMT Init + GenerateFuncDefinition ==================== |
| @@ -714,7 +714,8 @@ void BuildOutputPostChain(af::AscGraph &graph, OutputPostTopology topology, cons | |||
| 714 | } | 714 | } |
| 715 | 715 | ||
| 716 | af::AscGraph BuildPostReduceGraph(const std::string &suffix, bool reduce_outer = false, | 716 | af::AscGraph BuildPostReduceGraph(const std::string &suffix, bool reduce_outer = false, |
| 717 | - OutputPostTopology topology = OutputPostTopology::kSum) { | 717 | + OutputPostTopology topology = OutputPostTopology::kSum, |
| 718 | + bool cast_reduce_output = false, bool invalid_reduce_successor = false) { | ||
| 718 | af::AscGraph graph("indirect_load_post_reduce_ut_graph"); | 719 | af::AscGraph graph("indirect_load_post_reduce_ut_graph"); |
| 719 | const af::Expression s0 = graph.CreateSizeVar(2); | 720 | const af::Expression s0 = graph.CreateSizeVar(2); |
| 720 | const af::Expression s1 = suffix[0] == 'B' ? af::Expression(af::sym::kSymbolOne) : graph.CreateSizeVar(3); | 721 | const af::Expression s1 = suffix[0] == 'B' ? af::Expression(af::sym::kSymbolOne) : graph.CreateSizeVar(3); |
| @@ -786,7 +787,15 @@ af::AscGraph BuildPostReduceGraph(const std::string &suffix, bool reduce_outer = | |||
| 786 | sum.attr.api.compute_type = af::ComputeType::kComputeReduce; | 787 | sum.attr.api.compute_type = af::ComputeType::kComputeReduce; |
| 787 | sum.attr.sched.axis = output_axes; | 788 | sum.attr.sched.axis = output_axes; |
| 788 | SetNodeView(sum, post_dtype, output_axes, reduce_repeats, reduce_strides); | 789 | SetNodeView(sum, post_dtype, output_axes, reduce_repeats, reduce_strides); |
| 789 | - store.x = sum.y; | 790 | + if (cast_reduce_output) { |
| 791 | + // Model the Cast inserted by DtypeConsistency between Reduce and Store. | ||
| 792 | + af::ascir_op::Cast reduce_cast("reduce_cast"); | ||
| 793 | + reduce_cast.x = sum.y; | ||
| 794 | + SetNodeView(reduce_cast, post_dtype, output_axes, reduce_repeats, reduce_strides); | ||
| 795 | + store.x = reduce_cast.y; | ||
| 796 | + } else { | ||
| 797 | + store.x = sum.y; | ||
| 798 | + } | ||
| 790 | } | 799 | } |
| 791 | const auto &store_repeats = direct_output ? output_repeats : reduce_repeats; | 800 | const auto &store_repeats = direct_output ? output_repeats : reduce_repeats; |
| 792 | const auto &store_strides = direct_output ? output_strides : reduce_strides; | 801 | const auto &store_strides = direct_output ? output_strides : reduce_strides; |
| @@ -795,6 +804,22 @@ af::AscGraph BuildPostReduceGraph(const std::string &suffix, bool reduce_outer = | |||
| 795 | output.x = store.y; | 804 | output.x = store.y; |
| 796 | output.ir_attr.SetIndex(0); | 805 | output.ir_attr.SetIndex(0); |
| 797 | SetNodeView(output, post_dtype, output_axes, store_repeats, store_strides); | 806 | SetNodeView(output, post_dtype, output_axes, store_repeats, store_strides); |
| 807 | + if (invalid_reduce_successor) { | ||
| 808 | + af::ascir_op::Abs invalid_successor("invalid_reduce_successor"); | ||
| 809 | + invalid_successor.x = sum.y; | ||
| 810 | + SetNodeView(invalid_successor, post_dtype, output_axes, reduce_repeats, reduce_strides); | ||
| 811 | + const auto sum_node = graph.FindNode("sum"); | ||
| 812 | + const auto store_node = graph.FindNode("store"); | ||
| 813 | + const auto invalid_node = graph.FindNode("invalid_reduce_successor"); | ||
| 814 | + EXPECT_NE(sum_node, nullptr); | ||
| 815 | + EXPECT_NE(store_node, nullptr); | ||
| 816 | + EXPECT_NE(invalid_node, nullptr); | ||
| 817 | + if (sum_node != nullptr && store_node != nullptr && invalid_node != nullptr) { | ||
| 818 | + EXPECT_EQ(af::GraphUtils::ReplaceEdgeSrc(sum_node->GetOutDataAnchor(0), store_node->GetInDataAnchor(0), | ||
| 819 | + invalid_node->GetOutDataAnchor(0)), | ||
| 820 | + af::GRAPH_SUCCESS); | ||
| 821 | + } | ||
| 822 | + } | ||
| 798 | return graph; | 823 | return graph; |
| 799 | } | 824 | } |
| 800 | 825 | ||
| @@ -946,6 +971,15 @@ TEST(IndirectLoadScheduleCaseGeneratorTest, RejectsMixedZeroStrideAndPhysicalGap | |||
| 946 | EXPECT_EQ(layout.kind, ascgen_utils::indirect_load::IndirectLoadLayoutKind::kUnsupported); | 971 | EXPECT_EQ(layout.kind, ascgen_utils::indirect_load::IndirectLoadLayoutKind::kUnsupported); |
| 947 | } | 972 | } |
| 948 | 973 | ||
| 974 | +TEST(IndirectLoadScheduleCaseGeneratorTest, ClassifiesMixedZeroStrideAndPhysicalGapLayoutForSimt) { | ||
| 975 | + const ascgen_utils::indirect_load::LogicalTensorView view = { | ||
| 976 | + {0, 1, 2}, {af::Symbol(2), af::Symbol(3), af::Symbol(4)}, {af::Symbol(8), af::Symbol(0), af::Symbol(1)}}; | ||
| 977 | + ascgen_utils::indirect_load::IndirectLoadTensorLayout layout; | ||
| 978 | + ASSERT_EQ(ascgen_utils::indirect_load::ClassifyIndirectLoadLayout(view, layout, true), af::SUCCESS); | ||
| 979 | + EXPECT_EQ(layout.kind, ascgen_utils::indirect_load::IndirectLoadLayoutKind::kStrided); | ||
| 980 | + EXPECT_EQ(layout.physical_repeats, view.sizes); | ||
| 981 | +} | ||
| 982 | + | ||
| 949 | TEST(IndirectLoadScheduleCaseGeneratorTest, RejectsInvalidLogicalLayout) { | 983 | TEST(IndirectLoadScheduleCaseGeneratorTest, RejectsInvalidLogicalLayout) { |
| 950 | ascgen_utils::indirect_load::IndirectLoadTensorLayout layout; | 984 | ascgen_utils::indirect_load::IndirectLoadTensorLayout layout; |
| 951 | const ascgen_utils::indirect_load::LogicalTensorView rank_mismatch = { | 985 | const ascgen_utils::indirect_load::LogicalTensorView rank_mismatch = { |
| @@ -1504,6 +1538,27 @@ TEST(IndirectLoadScheduleCaseGeneratorTest, PostReduceRejectsMultipleReduceSegme | |||
| 1504 | EXPECT_NE(FindGeneratedGraphByTemplate(graphs, ascir::TemplateId::kIndirectLoadSK), graphs.end()); | 1538 | EXPECT_NE(FindGeneratedGraphByTemplate(graphs, ascir::TemplateId::kIndirectLoadSK), graphs.end()); |
| 1505 | } | 1539 | } |
| 1506 | 1540 | ||
| 1541 | +TEST(IndirectLoadScheduleCaseGeneratorTest, PostReduceAllowsDtypeCastBeforeStore) { | ||
| 1542 | + auto graph = BuildPostReduceGraph("R", false, OutputPostTopology::kSum, true); | ||
| 1543 | + optimize::IndirectLoadScheduleCaseGenerator generator; | ||
| 1544 | + std::vector<af::AscGraph> graphs; | ||
| 1545 | + std::vector<std::string> score_functions; | ||
| 1546 | + | ||
| 1547 | + ASSERT_EQ(generator.Generate(graph, graphs, score_functions), af::SUCCESS); | ||
| 1548 | + EXPECT_NE(FindGeneratedGraphByTemplate(graphs, ascir::TemplateId::kIndirectLoadSimt), graphs.end()); | ||
| 1549 | +} | ||
| 1550 | + | ||
| 1551 | +TEST(IndirectLoadScheduleCaseGeneratorTest, PostReduceRejectsNonCastSuccessor) { | ||
| 1552 | + auto graph = BuildPostReduceGraph("R", false, OutputPostTopology::kSum, false, true); | ||
| 1553 | + | ||
| 1554 | + optimize::IndirectLoadScheduleCaseGenerator generator; | ||
| 1555 | + std::vector<af::AscGraph> graphs; | ||
| 1556 | + std::vector<std::string> score_functions; | ||
| 1557 | + EXPECT_NE(generator.Generate(graph, graphs, score_functions), af::SUCCESS); | ||
| 1558 | + EXPECT_TRUE(graphs.empty()); | ||
| 1559 | + EXPECT_TRUE(score_functions.empty()); | ||
| 1560 | +} | ||
| 1561 | + | ||
| 1507 | TEST(IndirectLoadScheduleCaseGeneratorTest, PostReduceSkipsCommonZeroStrideAxes) { | 1562 | TEST(IndirectLoadScheduleCaseGeneratorTest, PostReduceSkipsCommonZeroStrideAxes) { |
| 1508 | for (const std::string suffix : {"RBR", "BRA"}) { | 1563 | for (const std::string suffix : {"RBR", "BRA"}) { |
| 1509 | auto graph = BuildPostReduceGraph(suffix); | 1564 | auto graph = BuildPostReduceGraph(suffix); |
| @@ -3,6 +3,10 @@ function(mark_indirect_load_codegen_and_e2e test_name) | |||
| 3 | target_compile_definitions(${test_name}_e2e_v2 PRIVATE ${ARGN}) | 3 | target_compile_definitions(${test_name}_e2e_v2 PRIVATE ${ARGN}) |
| 4 | endfunction() | 4 | endfunction() |
| 5 | 5 | ||
| 6 | +function(mark_indirect_load_red test_name) | ||
| 7 | + set_property(TEST ${test_name}_codegen_v2 ${test_name}_e2e_v2 APPEND PROPERTY LABELS red) | ||
| 8 | +endfunction() | ||
| 9 | + | ||
| 6 | function(mark_indirect_load_codegen test_name) | 10 | function(mark_indirect_load_codegen test_name) |
| 7 | target_compile_definitions(${test_name}_codegen_v2 PRIVATE ${ARGN}) | 11 | target_compile_definitions(${test_name}_codegen_v2 PRIVATE ${ARGN}) |
| 8 | endfunction() | 12 | endfunction() |
| @@ -349,6 +353,181 @@ do_backend_e2e_st_test(indirect_load_broadcast_index_where_simt_test | |||
| 349 | TEST_SRC test_e2e_indirect_load_store_kernel.cpp) | 353 | TEST_SRC test_e2e_indirect_load_store_kernel.cpp) |
| 350 | mark_indirect_load_codegen_and_e2e(indirect_load_broadcast_index_where_simt_test IL_CASE_BROADCAST_WHERE) | 354 | mark_indirect_load_codegen_and_e2e(indirect_load_broadcast_index_where_simt_test IL_CASE_BROADCAST_WHERE) |
| 351 | 355 | ||
| 356 | +# Strict regression for the user-provided GraphHint: Where + Broadcast + IndirectLoad | ||
| 357 | +# followed by ReduceSum, including the non-overlapping zero-stride/physical-gap table view. | ||
| 358 | +set(indirect_load_graph_hint_reduce_simt_test_workdir | ||
| 359 | + ${CMAKE_CURRENT_BINARY_DIR}/indirect_load_graph_hint_reduce_simt_test) | ||
| 360 | +file(MAKE_DIRECTORY ${indirect_load_graph_hint_reduce_simt_test_workdir}) | ||
| 361 | +do_backend_e2e_st_test(indirect_load_graph_hint_reduce_simt_test | ||
| 362 | + WORKDIR ${indirect_load_graph_hint_reduce_simt_test_workdir} | ||
| 363 | + CODEGEN indirect_load_store_backend_generator.cpp | ||
| 364 | + TILING_KEY 1 | ||
| 365 | + KERNEL_SRC | ||
| 366 | + indirect_load_graph_hint_reduce_simt_test_kernel.cpp | ||
| 367 | + indirect_load_graph_hint_reduce_simt_test_tiling.cpp | ||
| 368 | + autofuse_tiling_data.h | ||
| 369 | + TEST_SRC test_e2e_indirect_load_store_kernel.cpp) | ||
| 370 | +mark_indirect_load_codegen_and_e2e(indirect_load_graph_hint_reduce_simt_test | ||
| 371 | + IL_CASE_BROADCAST_WHERE IL_GRAPH_HINT_REDUCE) | ||
| 372 | + | ||
| 373 | +# Exact 30x3x23 GraphHint reproduction for the SIMD IndirectLoad gather path. | ||
| 374 | +set(indirect_load_graph_hint_simd_repro_workdir | ||
| 375 | + ${CMAKE_CURRENT_BINARY_DIR}/indirect_load_graph_hint_simd_repro) | ||
| 376 | +file(MAKE_DIRECTORY ${indirect_load_graph_hint_simd_repro_workdir}) | ||
| 377 | +do_backend_e2e_st_test(indirect_load_graph_hint_simd_repro | ||
| 378 | + WORKDIR ${indirect_load_graph_hint_simd_repro_workdir} | ||
| 379 | + CODEGEN indirect_load_store_backend_generator.cpp | ||
| 380 | + TILING_KEY 0 | ||
| 381 | + KERNEL_SRC | ||
| 382 | + indirect_load_graph_hint_simd_repro_kernel.cpp | ||
| 383 | + indirect_load_graph_hint_simd_repro_tiling.cpp | ||
| 384 | + autofuse_tiling_data.h | ||
| 385 | + TEST_SRC test_e2e_indirect_load_store_kernel.cpp) | ||
| 386 | +mark_indirect_load_codegen_and_e2e(indirect_load_graph_hint_simd_repro | ||
| 387 | + IL_CASE_BROADCAST_WHERE IL_GRAPH_HINT_SIMD_REPRO) | ||
| 388 | + | ||
| 389 | +# User graph: embedding + Sum over the lookup axis. | ||
| 390 | +set(indirect_load_user_embedding_sum_workdir | ||
| 391 | + ${CMAKE_CURRENT_BINARY_DIR}/indirect_load_user_embedding_sum) | ||
| 392 | +file(MAKE_DIRECTORY ${indirect_load_user_embedding_sum_workdir}) | ||
| 393 | +do_backend_e2e_st_test(indirect_load_user_embedding_sum | ||
| 394 | + WORKDIR ${indirect_load_user_embedding_sum_workdir} | ||
| 395 | + CODEGEN indirect_load_store_backend_generator.cpp | ||
| 396 | + TILING_KEY 1 | ||
| 397 | + KERNEL_SRC | ||
| 398 | + indirect_load_user_embedding_sum_kernel.cpp | ||
| 399 | + indirect_load_user_embedding_sum_tiling.cpp | ||
| 400 | + autofuse_tiling_data.h | ||
| 401 | + TEST_SRC test_e2e_indirect_load_store_kernel.cpp) | ||
| 402 | +mark_indirect_load_codegen_and_e2e(indirect_load_user_embedding_sum IL_CASE_BROADCAST_WHERE IL_USER_EMBEDDING_SUM) | ||
| 403 | + | ||
| 404 | +# User graph: embedding followed by scalar Mul. | ||
| 405 | +set(indirect_load_user_embedding_mul_workdir | ||
| 406 | + ${CMAKE_CURRENT_BINARY_DIR}/indirect_load_user_embedding_mul) | ||
| 407 | +file(MAKE_DIRECTORY ${indirect_load_user_embedding_mul_workdir}) | ||
| 408 | +do_backend_e2e_st_test(indirect_load_user_embedding_mul | ||
| 409 | + WORKDIR ${indirect_load_user_embedding_mul_workdir} | ||
| 410 | + CODEGEN indirect_load_store_backend_generator.cpp | ||
| 411 | + TILING_KEY 1 | ||
| 412 | + KERNEL_SRC | ||
| 413 | + indirect_load_user_embedding_mul_kernel.cpp | ||
| 414 | + indirect_load_user_embedding_mul_tiling.cpp | ||
| 415 | + autofuse_tiling_data.h | ||
| 416 | + TEST_SRC test_e2e_indirect_load_store_kernel.cpp) | ||
| 417 | +mark_indirect_load_codegen_and_e2e(indirect_load_user_embedding_mul IL_CASE_BROADCAST_WHERE IL_USER_EMBEDDING_MUL) | ||
| 418 | + | ||
| 419 | +# User graph: bf16 embedding with mean/rsqrt normalization and two outputs. | ||
| 420 | +set(indirect_load_user_layernorm_workdir | ||
| 421 | + ${CMAKE_CURRENT_BINARY_DIR}/indirect_load_user_layernorm) | ||
| 422 | +file(MAKE_DIRECTORY ${indirect_load_user_layernorm_workdir}) | ||
| 423 | +do_backend_e2e_st_test(indirect_load_user_layernorm | ||
| 424 | + WORKDIR ${indirect_load_user_layernorm_workdir} | ||
| 425 | + CODEGEN indirect_load_store_backend_generator.cpp | ||
| 426 | + TILING_KEY 1 | ||
| 427 | + KERNEL_SRC | ||
| 428 | + indirect_load_user_layernorm_kernel.cpp | ||
| 429 | + indirect_load_user_layernorm_tiling.cpp | ||
| 430 | + autofuse_tiling_data.h | ||
| 431 | + TEST_SRC test_e2e_indirect_load_store_kernel.cpp) | ||
| 432 | +mark_indirect_load_codegen_and_e2e(indirect_load_user_layernorm IL_CASE_BROADCAST_WHERE IL_USER_LAYERNORM) | ||
| 433 | + | ||
| 434 | +# Same LayerNorm split topology with a small shape; force the SIMD candidate to | ||
| 435 | +# determine whether the branch itself, rather than UB pressure, is the blocker. | ||
| 436 | +set(indirect_load_user_layernorm_simd_workdir | ||
| 437 | + ${CMAKE_CURRENT_BINARY_DIR}/indirect_load_user_layernorm_simd) | ||
| 438 | +file(MAKE_DIRECTORY ${indirect_load_user_layernorm_simd_workdir}) | ||
| 439 | +do_backend_e2e_st_test(indirect_load_user_layernorm_simd | ||
| 440 | + WORKDIR ${indirect_load_user_layernorm_simd_workdir} | ||
| 441 | + CODEGEN indirect_load_store_backend_generator.cpp | ||
| 442 | + TILING_KEY 1 | ||
| 443 | + KERNEL_SRC | ||
| 444 | + indirect_load_user_layernorm_simd_kernel.cpp | ||
| 445 | + indirect_load_user_layernorm_simd_tiling.cpp | ||
| 446 | + autofuse_tiling_data.h | ||
| 447 | + TEST_SRC test_e2e_indirect_load_store_kernel.cpp) | ||
| 448 | +mark_indirect_load_codegen_and_e2e(indirect_load_user_layernorm_simd | ||
| 449 | + IL_CASE_BROADCAST_WHERE IL_USER_LAYERNORM IL_USER_LAYERNORM_SIMD) | ||
| 450 | + | ||
| 451 | +# User graph: one embedding result fans out to Exp and Abs, then merges by Add. | ||
| 452 | +# Run both IL SIMD and IL SIMT templates to compare branch handling. | ||
| 453 | +set(indirect_load_user_embedding_exp_abs_add_simd_workdir | ||
| 454 | + ${CMAKE_CURRENT_BINARY_DIR}/indirect_load_user_embedding_exp_abs_add_simd) | ||
| 455 | +file(MAKE_DIRECTORY ${indirect_load_user_embedding_exp_abs_add_simd_workdir}) | ||
| 456 | +do_backend_e2e_st_test(indirect_load_user_embedding_exp_abs_add_simd | ||
| 457 | + WORKDIR ${indirect_load_user_embedding_exp_abs_add_simd_workdir} | ||
| 458 | + CODEGEN indirect_load_store_backend_generator.cpp | ||
| 459 | + TILING_KEY 1 | ||
| 460 | + KERNEL_SRC | ||
| 461 | + indirect_load_user_embedding_exp_abs_add_simd_kernel.cpp | ||
| 462 | + indirect_load_user_embedding_exp_abs_add_simd_tiling.cpp | ||
| 463 | + autofuse_tiling_data.h | ||
| 464 | + TEST_SRC test_e2e_indirect_load_store_kernel.cpp) | ||
| 465 | +mark_indirect_load_codegen_and_e2e(indirect_load_user_embedding_exp_abs_add_simd | ||
| 466 | + IL_USER_EMBEDDING_EXP_ABS_ADD IL_USER_EMBEDDING_EXP_ABS_ADD_SIMD) | ||
| 467 | + | ||
| 468 | +set(indirect_load_user_embedding_exp_abs_add_simt_workdir | ||
| 469 | + ${CMAKE_CURRENT_BINARY_DIR}/indirect_load_user_embedding_exp_abs_add_simt) | ||
| 470 | +file(MAKE_DIRECTORY ${indirect_load_user_embedding_exp_abs_add_simt_workdir}) | ||
| 471 | +do_backend_e2e_st_test(indirect_load_user_embedding_exp_abs_add_simt | ||
| 472 | + WORKDIR ${indirect_load_user_embedding_exp_abs_add_simt_workdir} | ||
| 473 | + CODEGEN indirect_load_store_backend_generator.cpp | ||
| 474 | + TILING_KEY 1 | ||
| 475 | + KERNEL_SRC | ||
| 476 | + indirect_load_user_embedding_exp_abs_add_simt_kernel.cpp | ||
| 477 | + indirect_load_user_embedding_exp_abs_add_simt_tiling.cpp | ||
| 478 | + autofuse_tiling_data.h | ||
| 479 | + TEST_SRC test_e2e_indirect_load_store_kernel.cpp) | ||
| 480 | +mark_indirect_load_codegen_and_e2e(indirect_load_user_embedding_exp_abs_add_simt | ||
| 481 | + IL_USER_EMBEDDING_EXP_ABS_ADD) | ||
| 482 | + | ||
| 483 | +# Fan-out matrix: split at IndirectLoad or at a post-IndirectLoad elementwise node, | ||
| 484 | +# with either two ordinary Stores or one ordinary Store plus one Reduce Store. | ||
| 485 | +function(add_indirect_load_user_fanout_case test_name) | ||
| 486 | + set(workdir ${CMAKE_CURRENT_BINARY_DIR}/${test_name}) | ||
| 487 | + file(MAKE_DIRECTORY ${workdir}) | ||
| 488 | + do_backend_e2e_st_test(${test_name} | ||
| 489 | + WORKDIR ${workdir} | ||
| 490 | + CODEGEN indirect_load_store_backend_generator.cpp | ||
| 491 | + TILING_KEY 1 | ||
| 492 | + KERNEL_SRC | ||
| 493 | + ${test_name}_kernel.cpp | ||
| 494 | + ${test_name}_tiling.cpp | ||
| 495 | + autofuse_tiling_data.h | ||
| 496 | + TEST_SRC test_e2e_indirect_load_store_kernel.cpp) | ||
| 497 | + mark_indirect_load_codegen_and_e2e(${test_name} IL_USER_FANOUT ${ARGN}) | ||
| 498 | +endfunction() | ||
| 499 | + | ||
| 500 | +add_indirect_load_user_fanout_case(indirect_load_user_fanout_direct_stores_simd | ||
| 501 | + IL_USER_FANOUT_SIMD) | ||
| 502 | +add_indirect_load_user_fanout_case(indirect_load_user_fanout_direct_stores_simt) | ||
| 503 | +add_indirect_load_user_fanout_case(indirect_load_user_fanout_direct_reduce_simd | ||
| 504 | + IL_USER_FANOUT_SIMD IL_USER_FANOUT_REDUCE) | ||
| 505 | +add_indirect_load_user_fanout_case(indirect_load_user_fanout_direct_reduce_simt | ||
| 506 | + IL_USER_FANOUT_REDUCE) | ||
| 507 | +add_indirect_load_user_fanout_case(indirect_load_user_fanout_post_stores_simd | ||
| 508 | + IL_USER_FANOUT_SIMD IL_USER_FANOUT_POST) | ||
| 509 | +add_indirect_load_user_fanout_case(indirect_load_user_fanout_post_stores_simt | ||
| 510 | + IL_USER_FANOUT_POST) | ||
| 511 | +add_indirect_load_user_fanout_case(indirect_load_user_fanout_post_reduce_simd | ||
| 512 | + IL_USER_FANOUT_SIMD IL_USER_FANOUT_POST IL_USER_FANOUT_REDUCE) | ||
| 513 | +add_indirect_load_user_fanout_case(indirect_load_user_fanout_post_reduce_simt | ||
| 514 | + IL_USER_FANOUT_POST IL_USER_FANOUT_REDUCE) | ||
| 515 | + | ||
| 516 | +# User graph: two IndirectLoads over two tensors followed by Add (dual-IL graph). | ||
| 517 | +set(indirect_load_user_add_gather_workdir | ||
| 518 | + ${CMAKE_CURRENT_BINARY_DIR}/indirect_load_user_add_gather) | ||
| 519 | +file(MAKE_DIRECTORY ${indirect_load_user_add_gather_workdir}) | ||
| 520 | +do_backend_e2e_st_test(indirect_load_user_add_gather | ||
| 521 | + WORKDIR ${indirect_load_user_add_gather_workdir} | ||
| 522 | + CODEGEN indirect_load_store_backend_generator.cpp | ||
| 523 | + TILING_KEY 1 | ||
| 524 | + KERNEL_SRC | ||
| 525 | + indirect_load_user_add_gather_kernel.cpp | ||
| 526 | + indirect_load_user_add_gather_tiling.cpp | ||
| 527 | + autofuse_tiling_data.h | ||
| 528 | + TEST_SRC test_e2e_indirect_load_store_kernel.cpp) | ||
| 529 | +mark_indirect_load_codegen_and_e2e(indirect_load_user_add_gather IL_CASE_BROADCAST_WHERE IL_DUAL_IL_GATHER) | ||
| 530 | + | ||
| 352 | # Regression: SIMT post-reduce Output() evaluator uses full output_index instead of address.index_offset | 531 | # Regression: SIMT post-reduce Output() evaluator uses full output_index instead of address.index_offset |
| 353 | # to access index-tensor GM, causing OOB reads when output_index exceeds index tensor size. | 532 | # to access index-tensor GM, causing OOB reads when output_index exceeds index tensor size. |
| 354 | set(indirect_load_embedding_reduce_simt_test_workdir | 533 | set(indirect_load_embedding_reduce_simt_test_workdir |
| @@ -24,6 +24,27 @@ | |||
| 24 | extern "C" int64_t AutofuseTiling(AutofuseTilingData *, uint32_t *, uint32_t *, uint32_t, uint32_t); | 24 | extern "C" int64_t AutofuseTiling(AutofuseTilingData *, uint32_t *, uint32_t *, uint32_t, uint32_t); |
| 25 | 25 | ||
| 26 | 26 | ||
| 27 | + | ||
| 28 | +extern "C" __global__ __aicore__ void user_fanout(GM_ADDR indices, GM_ADDR embedding, GM_ADDR weight, GM_ADDR output0, | ||
| 29 | + GM_ADDR output1, GM_ADDR workspace, GM_ADDR gm_tiling_data); | ||
| 30 | + | ||
| 31 | +extern "C" __global__ __aicore__ void user_embedding_exp_abs_add(GM_ADDR indices, GM_ADDR embedding, GM_ADDR output, | ||
| 32 | + GM_ADDR workspace, GM_ADDR gm_tiling_data); | ||
| 33 | + | ||
| 34 | +extern "C" __global__ __aicore__ void user_embedding_sum(GM_ADDR table, GM_ADDR indices, GM_ADDR output, | ||
| 35 | + GM_ADDR workspace, GM_ADDR gm_tiling_data); | ||
| 36 | + | ||
| 37 | +extern "C" __global__ __aicore__ void user_embedding_mul(GM_ADDR table, GM_ADDR indices, float scale, GM_ADDR output, | ||
| 38 | + GM_ADDR workspace, GM_ADDR gm_tiling_data); | ||
| 39 | + | ||
| 40 | +extern "C" __global__ __aicore__ void user_layernorm(GM_ADDR indices, GM_ADDR embedding, GM_ADDR weight, | ||
| 41 | + GM_ADDR raw_output, GM_ADDR square_output, GM_ADDR workspace, | ||
| 42 | + GM_ADDR gm_tiling_data); | ||
| 43 | + | ||
| 44 | +extern "C" __global__ __aicore__ void user_add_gather(GM_ADDR input0, GM_ADDR input1, GM_ADDR indices, GM_ADDR output, | ||
| 45 | + GM_ADDR workspace, GM_ADDR gm_tiling_data); | ||
| 46 | + | ||
| 47 | + | ||
| 27 | namespace indirect_load_test { | 48 | namespace indirect_load_test { |
| 28 | inline void GmFree(void *ptr) { | 49 | inline void GmFree(void *ptr) { |
| 29 | AscendC::GmFree(ptr); | 50 | AscendC::GmFree(ptr); |
| @@ -840,7 +861,336 @@ TEST(E2EIndirectLoadBroadcast, GeneratedKernelMatchesReference) { | |||
| 840 | 861 | ||
| 841 | 862 | ||
| 842 | 863 | ||
| 843 | -#if defined(IL_CASE_BROADCAST_WHERE) | 864 | +#if defined(IL_USER_FANOUT) |
| 865 | +TEST(UserGraphConstruction, GeneratedKernelMatchesReference) { | ||
| 866 | + constexpr int32_t kRows = 2; | ||
| 867 | + constexpr int32_t kDim = 16; | ||
| 868 | + constexpr int32_t kTableRows = 32; | ||
| 869 | + auto *indices = static_cast<int64_t *>(AscendC::GmAlloc(sizeof(int64_t) * kRows)); | ||
| 870 | + auto *embedding = static_cast<bfloat16_t *>(AscendC::GmAlloc(sizeof(bfloat16_t) * kTableRows * kDim)); | ||
| 871 | + auto *weight = static_cast<bfloat16_t *>(AscendC::GmAlloc(sizeof(bfloat16_t) * kRows * kDim)); | ||
| 872 | + auto *output0 = static_cast<bfloat16_t *>(AscendC::GmAlloc(sizeof(bfloat16_t) * kRows * kDim)); | ||
| 873 | + | ||
| 874 | + auto *output1 = static_cast<bfloat16_t *>(AscendC::GmAlloc(sizeof(bfloat16_t) * kRows)); | ||
| 875 | + | ||
| 876 | + auto *output1 = static_cast<bfloat16_t *>(AscendC::GmAlloc(sizeof(bfloat16_t) * kRows * kDim)); | ||
| 877 | + | ||
| 878 | + ASSERT_NE(indices, nullptr); | ||
| 879 | + ASSERT_NE(embedding, nullptr); | ||
| 880 | + ASSERT_NE(weight, nullptr); | ||
| 881 | + ASSERT_NE(output0, nullptr); | ||
| 882 | + ASSERT_NE(output1, nullptr); | ||
| 883 | + for (int32_t row = 0; row < kRows; ++row) { | ||
| 884 | + indices[row] = row + 1; | ||
| 885 | + } | ||
| 886 | + for (int32_t row = 0; row < kTableRows; ++row) { | ||
| 887 | + for (int32_t col = 0; col < kDim; ++col) { | ||
| 888 | + embedding[row * kDim + col] = static_cast<bfloat16_t>((row + 1) * 0.01F + col * 0.001F); | ||
| 889 | + } | ||
| 890 | + } | ||
| 891 | + std::fill_n(weight, kRows * kDim, static_cast<bfloat16_t>(0.0F)); | ||
| 892 | + std::fill_n(output0, kRows * kDim, static_cast<bfloat16_t>(0.0F)); | ||
| 893 | + | ||
| 894 | + std::fill_n(output1, kRows, static_cast<bfloat16_t>(0.0F)); | ||
| 895 | + | ||
| 896 | + std::fill_n(output1, kRows * kDim, static_cast<bfloat16_t>(0.0F)); | ||
| 897 | + | ||
| 898 | + | ||
| 899 | + AutofuseTilingData tiling_data{}; | ||
| 900 | + uint32_t workspace_size = 0U; | ||
| 901 | + uint32_t block_dim = 48U; | ||
| 902 | + ASSERT_EQ(AutofuseTiling(&tiling_data, &workspace_size, &block_dim, 48U, 192U * 1024U), 0); | ||
| 903 | + void *workspace = workspace_size == 0U ? nullptr : AscendC::GmAlloc(workspace_size); | ||
| 904 | + ASSERT_TRUE(workspace_size == 0U || workspace != nullptr); | ||
| 905 | + AscendC::SetKernelMode(KernelMode::AIV_MODE); | ||
| 906 | + ICPU_RUN_KF(user_fanout, block_dim, reinterpret_cast<uint8_t *>(indices), reinterpret_cast<uint8_t *>(embedding), | ||
| 907 | + reinterpret_cast<uint8_t *>(weight), reinterpret_cast<uint8_t *>(output0), | ||
| 908 | + reinterpret_cast<uint8_t *>(output1), reinterpret_cast<uint8_t *>(workspace), | ||
| 909 | + reinterpret_cast<uint8_t *>(&tiling_data)); | ||
| 910 | + | ||
| 911 | + for (int32_t row = 0; row < kRows; ++row) { | ||
| 912 | + | ||
| 913 | + float expected_reduce = 0.0F; | ||
| 914 | + | ||
| 915 | + for (int32_t col = 0; col < kDim; ++col) { | ||
| 916 | + const float value = static_cast<float>(embedding[indices[row] * kDim + col]); | ||
| 917 | + | ||
| 918 | + const float source = std::fabs(value); | ||
| 919 | + | ||
| 920 | + const float source = value; | ||
| 921 | + | ||
| 922 | + EXPECT_NEAR(static_cast<float>(output0[row * kDim + col]), std::exp(source), 0.125F) | ||
| 923 | + << "output0 row=" << row << ", col=" << col; | ||
| 924 | + | ||
| 925 | + expected_reduce += source * source; | ||
| 926 | + | ||
| 927 | + EXPECT_NEAR(static_cast<float>(output1[row * kDim + col]), std::fabs(source), 0.125F) | ||
| 928 | + << "output1 row=" << row << ", col=" << col; | ||
| 929 | + | ||
| 930 | + } | ||
| 931 | + | ||
| 932 | + EXPECT_NEAR(static_cast<float>(output1[row]), expected_reduce, 0.5F) << "output1 row=" << row; | ||
| 933 | + | ||
| 934 | + } | ||
| 935 | + if (workspace != nullptr) AscendC::GmFree(workspace); | ||
| 936 | + AscendC::GmFree(indices); | ||
| 937 | + AscendC::GmFree(embedding); | ||
| 938 | + AscendC::GmFree(weight); | ||
| 939 | + AscendC::GmFree(output0); | ||
| 940 | + AscendC::GmFree(output1); | ||
| 941 | +} | ||
| 942 | + | ||
| 943 | + | ||
| 944 | +TEST(UserGraphConstruction, GeneratedKernelMatchesReference) { | ||
| 945 | + constexpr int32_t kRows = 8; | ||
| 946 | + constexpr int32_t kDim = 16; | ||
| 947 | + constexpr int32_t kTableRows = 100; | ||
| 948 | + constexpr int32_t kTableElements = kTableRows * kDim; | ||
| 949 | + auto *indices = static_cast<int64_t *>(AscendC::GmAlloc(sizeof(int64_t) * kRows)); | ||
| 950 | + auto *embedding = static_cast<float *>(AscendC::GmAlloc(sizeof(float) * kTableElements)); | ||
| 951 | + auto *output = static_cast<float *>(AscendC::GmAlloc(sizeof(float) * kRows * kDim)); | ||
| 952 | + ASSERT_NE(indices, nullptr); | ||
| 953 | + ASSERT_NE(embedding, nullptr); | ||
| 954 | + ASSERT_NE(output, nullptr); | ||
| 955 | + | ||
| 956 | + for (int32_t row = 0; row < kRows; ++row) { | ||
| 957 | + indices[row] = static_cast<int64_t>(row + 1); | ||
| 958 | + } | ||
| 959 | + for (int32_t row = 0; row < kTableRows; ++row) { | ||
| 960 | + for (int32_t col = 0; col < kDim; ++col) { | ||
| 961 | + embedding[row * kDim + col] = static_cast<float>(row - 40) * 0.01F + static_cast<float>(col) * 0.001F; | ||
| 962 | + } | ||
| 963 | + } | ||
| 964 | + std::fill_n(output, kRows * kDim, 0.0F); | ||
| 965 | + | ||
| 966 | + AutofuseTilingData tiling_data{}; | ||
| 967 | + uint32_t workspace_size = 0U; | ||
| 968 | + uint32_t block_dim = 48U; | ||
| 969 | + ASSERT_EQ(AutofuseTiling(&tiling_data, &workspace_size, &block_dim, 48U, 192U * 1024U), 0); | ||
| 970 | + void *workspace = workspace_size == 0U ? nullptr : AscendC::GmAlloc(workspace_size); | ||
| 971 | + ASSERT_TRUE(workspace_size == 0U || workspace != nullptr); | ||
| 972 | + | ||
| 973 | + AscendC::SetKernelMode(KernelMode::AIV_MODE); | ||
| 974 | + ICPU_RUN_KF(user_embedding_exp_abs_add, block_dim, reinterpret_cast<uint8_t *>(indices), | ||
| 975 | + reinterpret_cast<uint8_t *>(embedding), reinterpret_cast<uint8_t *>(output), | ||
| 976 | + reinterpret_cast<uint8_t *>(workspace), reinterpret_cast<uint8_t *>(&tiling_data)); | ||
| 977 | + | ||
| 978 | + for (int32_t row = 0; row < kRows; ++row) { | ||
| 979 | + const int64_t index = indices[row]; | ||
| 980 | + for (int32_t col = 0; col < kDim; ++col) { | ||
| 981 | + const float value = embedding[index * kDim + col]; | ||
| 982 | + const float expected = std::exp(value) + std::fabs(value); | ||
| 983 | + EXPECT_NEAR(output[row * kDim + col], expected, 1.0e-4F) << "row=" << row << ", col=" << col; | ||
| 984 | + } | ||
| 985 | + } | ||
| 986 | + | ||
| 987 | + if (workspace != nullptr) { | ||
| 988 | + AscendC::GmFree(workspace); | ||
| 989 | + } | ||
| 990 | + AscendC::GmFree(indices); | ||
| 991 | + AscendC::GmFree(embedding); | ||
| 992 | + AscendC::GmFree(output); | ||
| 993 | +} | ||
| 994 | + | ||
| 995 | + | ||
| 996 | +TEST(UserGraphConstruction, GeneratedKernelMatchesReference) { | ||
| 997 | + constexpr int32_t kRows = 4; | ||
| 998 | + constexpr int32_t kLookups = 4; | ||
| 999 | + constexpr int32_t kDim = 16; | ||
| 1000 | + constexpr int32_t kTableRows = 100; | ||
| 1001 | + constexpr int32_t kIndexStride = 8; | ||
| 1002 | + const int32_t kIndexStorage = (kRows - 1) * kIndexStride + kLookups; | ||
| 1003 | + auto *indices = static_cast<int32_t *>(AscendC::GmAlloc(sizeof(int32_t) * kIndexStorage)); | ||
| 1004 | + auto *table = static_cast<float *>(AscendC::GmAlloc(sizeof(float) * kTableRows * kDim)); | ||
| 1005 | + auto *output = static_cast<float *>(AscendC::GmAlloc(sizeof(float) * kRows * kDim)); | ||
| 1006 | + ASSERT_NE(indices, nullptr); | ||
| 1007 | + ASSERT_NE(table, nullptr); | ||
| 1008 | + ASSERT_NE(output, nullptr); | ||
| 1009 | + for (int32_t row = 0; row < kTableRows; ++row) { | ||
| 1010 | + for (int32_t col = 0; col < kDim; ++col) { | ||
| 1011 | + table[row * kDim + col] = static_cast<float>(row * kDim + col) * 0.01F; | ||
| 1012 | + } | ||
| 1013 | + } | ||
| 1014 | + for (int32_t row = 0; row < kRows; ++row) { | ||
| 1015 | + for (int32_t lookup = 0; lookup < kLookups; ++lookup) { | ||
| 1016 | + indices[row * kIndexStride + lookup] = (row * kLookups + lookup + 1) % kTableRows; | ||
| 1017 | + } | ||
| 1018 | + } | ||
| 1019 | + std::fill_n(output, kRows * kDim, 0.0F); | ||
| 1020 | + | ||
| 1021 | + AutofuseTilingData tiling_data{}; | ||
| 1022 | + tiling_data.set_ks0(kRows); | ||
| 1023 | + tiling_data.set_s44(kIndexStride); | ||
| 1024 | + uint32_t workspace_size = 0U; | ||
| 1025 | + uint32_t block_dim = 48U; | ||
| 1026 | + ASSERT_EQ(AutofuseTiling(&tiling_data, &workspace_size, &block_dim, 48U, 192U * 1024U), 0); | ||
| 1027 | + void *workspace = workspace_size == 0U ? nullptr : AscendC::GmAlloc(workspace_size); | ||
| 1028 | + ASSERT_TRUE(workspace_size == 0U || workspace != nullptr); | ||
| 1029 | + AscendC::SetKernelMode(KernelMode::AIV_MODE); | ||
| 1030 | + ICPU_RUN_KF(user_embedding_sum, block_dim, reinterpret_cast<uint8_t *>(table), reinterpret_cast<uint8_t *>(indices), | ||
| 1031 | + reinterpret_cast<uint8_t *>(output), reinterpret_cast<uint8_t *>(workspace), | ||
| 1032 | + reinterpret_cast<uint8_t *>(&tiling_data)); | ||
| 1033 | + for (int32_t row = 0; row < kRows; ++row) { | ||
| 1034 | + for (int32_t col = 0; col < kDim; ++col) { | ||
| 1035 | + float expected = 0.0F; | ||
| 1036 | + for (int32_t lookup = 0; lookup < kLookups; ++lookup) { | ||
| 1037 | + expected += table[indices[row * kIndexStride + lookup] * kDim + col]; | ||
| 1038 | + } | ||
| 1039 | + EXPECT_NEAR(output[row * kDim + col], expected, 0.0625F) << "row=" << row << ", col=" << col; | ||
| 1040 | + } | ||
| 1041 | + } | ||
| 1042 | + if (workspace != nullptr) AscendC::GmFree(workspace); | ||
| 1043 | + AscendC::GmFree(indices); | ||
| 1044 | + AscendC::GmFree(table); | ||
| 1045 | + AscendC::GmFree(output); | ||
| 1046 | +} | ||
| 1047 | + | ||
| 1048 | + | ||
| 1049 | +TEST(UserGraphConstruction, GeneratedKernelMatchesReference) { | ||
| 1050 | + constexpr int32_t kRows = 1024; | ||
| 1051 | + constexpr int32_t kDim = 2048; | ||
| 1052 | + constexpr int32_t kTableRows = 2; | ||
| 1053 | + auto *indices = static_cast<int64_t *>(AscendC::GmAlloc(sizeof(int64_t) * kRows)); | ||
| 1054 | + auto *table = static_cast<half *>(AscendC::GmAlloc(sizeof(half) * kTableRows * kDim)); | ||
| 1055 | + auto *output = static_cast<half *>(AscendC::GmAlloc(sizeof(half) * kRows * kDim)); | ||
| 1056 | + ASSERT_NE(indices, nullptr); | ||
| 1057 | + ASSERT_NE(table, nullptr); | ||
| 1058 | + ASSERT_NE(output, nullptr); | ||
| 1059 | + constexpr float scale = 0.5F; | ||
| 1060 | + for (int32_t row = 0; row < kRows; ++row) indices[row] = row % kTableRows; | ||
| 1061 | + for (int32_t row = 0; row < kTableRows; ++row) { | ||
| 1062 | + for (int32_t col = 0; col < kDim; ++col) { | ||
| 1063 | + table[row * kDim + col] = static_cast<half>((row + 1) * 0.01F + col * 0.001F); | ||
| 1064 | + } | ||
| 1065 | + } | ||
| 1066 | + std::fill_n(output, kRows * kDim, static_cast<half>(0.0F)); | ||
| 1067 | + AutofuseTilingData tiling_data{}; | ||
| 1068 | + uint32_t workspace_size = 0U; | ||
| 1069 | + uint32_t block_dim = 48U; | ||
| 1070 | + ASSERT_EQ(AutofuseTiling(&tiling_data, &workspace_size, &block_dim, 48U, 192U * 1024U), 0); | ||
| 1071 | + void *workspace = workspace_size == 0U ? nullptr : AscendC::GmAlloc(workspace_size); | ||
| 1072 | + ASSERT_TRUE(workspace_size == 0U || workspace != nullptr); | ||
| 1073 | + AscendC::SetKernelMode(KernelMode::AIV_MODE); | ||
| 1074 | + ICPU_RUN_KF(user_embedding_mul, block_dim, reinterpret_cast<uint8_t *>(table), reinterpret_cast<uint8_t *>(indices), | ||
| 1075 | + scale, reinterpret_cast<uint8_t *>(output), reinterpret_cast<uint8_t *>(workspace), | ||
| 1076 | + reinterpret_cast<uint8_t *>(&tiling_data)); | ||
| 1077 | + for (int32_t row = 0; row < kRows; ++row) { | ||
| 1078 | + for (int32_t col = 0; col < kDim; ++col) { | ||
| 1079 | + const float expected = static_cast<float>(table[indices[row] * kDim + col]) * scale; | ||
| 1080 | + EXPECT_NEAR(static_cast<float>(output[row * kDim + col]), expected, 0.0625F) << "row=" << row << ", col=" << col; | ||
| 1081 | + } | ||
| 1082 | + } | ||
| 1083 | + if (workspace != nullptr) AscendC::GmFree(workspace); | ||
| 1084 | + AscendC::GmFree(indices); | ||
| 1085 | + AscendC::GmFree(table); | ||
| 1086 | + AscendC::GmFree(output); | ||
| 1087 | +} | ||
| 1088 | + | ||
| 1089 | + | ||
| 1090 | +TEST(UserGraphConstruction, GeneratedKernelMatchesReference) { | ||
| 1091 | + | ||
| 1092 | + constexpr int32_t kRows = 2; | ||
| 1093 | + constexpr int32_t kDim = 16; | ||
| 1094 | + constexpr int32_t kTableRows = 100; | ||
| 1095 | + | ||
| 1096 | + constexpr int32_t kRows = 21; | ||
| 1097 | + constexpr int32_t kDim = 2048; | ||
| 1098 | + constexpr int32_t kTableRows = 102400; | ||
| 1099 | + | ||
| 1100 | + auto *indices = static_cast<int64_t *>(AscendC::GmAlloc(sizeof(int64_t) * kRows)); | ||
| 1101 | + auto *embedding = static_cast<bfloat16_t *>(AscendC::GmAlloc(sizeof(bfloat16_t) * kTableRows * kDim)); | ||
| 1102 | + auto *weight = static_cast<bfloat16_t *>(AscendC::GmAlloc(sizeof(bfloat16_t) * kRows * kDim)); | ||
| 1103 | + auto *raw_output = static_cast<bfloat16_t *>(AscendC::GmAlloc(sizeof(bfloat16_t) * kRows * kDim)); | ||
| 1104 | + auto *square_output = static_cast<bfloat16_t *>(AscendC::GmAlloc(sizeof(bfloat16_t) * kRows)); | ||
| 1105 | + ASSERT_NE(indices, nullptr); | ||
| 1106 | + ASSERT_NE(embedding, nullptr); | ||
| 1107 | + ASSERT_NE(weight, nullptr); | ||
| 1108 | + ASSERT_NE(raw_output, nullptr); | ||
| 1109 | + ASSERT_NE(square_output, nullptr); | ||
| 1110 | + for (int32_t row = 0; row < kRows; ++row) indices[row] = row % kTableRows; | ||
| 1111 | + for (int32_t row = 0; row < kTableRows; ++row) { | ||
| 1112 | + for (int32_t col = 0; col < kDim; ++col) { | ||
| 1113 | + embedding[row * kDim + col] = static_cast<bfloat16_t>((row + 1) * 0.01F + col * 0.001F); | ||
| 1114 | + } | ||
| 1115 | + } | ||
| 1116 | + std::fill_n(weight, kRows * kDim, static_cast<bfloat16_t>(0.0F)); | ||
| 1117 | + std::fill_n(raw_output, kRows * kDim, static_cast<bfloat16_t>(0.0F)); | ||
| 1118 | + std::fill_n(square_output, kRows, static_cast<bfloat16_t>(0.0F)); | ||
| 1119 | + AutofuseTilingData tiling_data{}; | ||
| 1120 | + uint32_t workspace_size = 0U; | ||
| 1121 | + uint32_t block_dim = 48U; | ||
| 1122 | + ASSERT_EQ(AutofuseTiling(&tiling_data, &workspace_size, &block_dim, 48U, 192U * 1024U), 0); | ||
| 1123 | + void *workspace = workspace_size == 0U ? nullptr : AscendC::GmAlloc(workspace_size); | ||
| 1124 | + ASSERT_TRUE(workspace_size == 0U || workspace != nullptr); | ||
| 1125 | + AscendC::SetKernelMode(KernelMode::AIV_MODE); | ||
| 1126 | + ICPU_RUN_KF(user_layernorm, block_dim, reinterpret_cast<uint8_t *>(indices), reinterpret_cast<uint8_t *>(embedding), | ||
| 1127 | + reinterpret_cast<uint8_t *>(weight), reinterpret_cast<uint8_t *>(raw_output), | ||
| 1128 | + reinterpret_cast<uint8_t *>(square_output), reinterpret_cast<uint8_t *>(workspace), | ||
| 1129 | + reinterpret_cast<uint8_t *>(&tiling_data)); | ||
| 1130 | + for (int32_t row = 0; row < kRows; ++row) { | ||
| 1131 | + float expected_square = 0.0F; | ||
| 1132 | + for (int32_t col = 0; col < kDim; ++col) { | ||
| 1133 | + const float expected_raw = static_cast<float>(embedding[indices[row] * kDim + col]); | ||
| 1134 | + expected_square += expected_raw * expected_raw; | ||
| 1135 | + EXPECT_NEAR(static_cast<float>(raw_output[row * kDim + col]), expected_raw, 0.0625F) | ||
| 1136 | + << "raw row=" << row << ", col=" << col; | ||
| 1137 | + } | ||
| 1138 | + EXPECT_NEAR(static_cast<float>(square_output[row]), expected_square, 0.5F) << "square row=" << row; | ||
| 1139 | + } | ||
| 1140 | + if (workspace != nullptr) AscendC::GmFree(workspace); | ||
| 1141 | + AscendC::GmFree(indices); | ||
| 1142 | + AscendC::GmFree(embedding); | ||
| 1143 | + AscendC::GmFree(weight); | ||
| 1144 | + AscendC::GmFree(raw_output); | ||
| 1145 | + AscendC::GmFree(square_output); | ||
| 1146 | +} | ||
| 1147 | + | ||
| 1148 | + | ||
| 1149 | +TEST(UserGraphConstruction, GeneratedKernelMatchesReference) { | ||
| 1150 | + constexpr int32_t kRows = 1024 * 1025; | ||
| 1151 | + constexpr int32_t kInputWidth = 10; | ||
| 1152 | + constexpr int32_t kOutputWidth = 5; | ||
| 1153 | + auto *input0 = static_cast<float *>(AscendC::GmAlloc(sizeof(float) * kRows * kInputWidth)); | ||
| 1154 | + auto *input1 = static_cast<float *>(AscendC::GmAlloc(sizeof(float) * kRows * kInputWidth)); | ||
| 1155 | + auto *indices = static_cast<int64_t *>(AscendC::GmAlloc(sizeof(int64_t) * kRows * kOutputWidth)); | ||
| 1156 | + auto *output = static_cast<float *>(AscendC::GmAlloc(sizeof(float) * kRows * kOutputWidth)); | ||
| 1157 | + ASSERT_NE(input0, nullptr); | ||
| 1158 | + ASSERT_NE(input1, nullptr); | ||
| 1159 | + ASSERT_NE(indices, nullptr); | ||
| 1160 | + ASSERT_NE(output, nullptr); | ||
| 1161 | + for (int32_t row = 0; row < kRows; ++row) { | ||
| 1162 | + for (int32_t col = 0; col < kInputWidth; ++col) { | ||
| 1163 | + input0[row * kInputWidth + col] = row * 0.001F + col; | ||
| 1164 | + input1[row * kInputWidth + col] = row * 0.002F - col; | ||
| 1165 | + } | ||
| 1166 | + for (int32_t col = 0; col < kOutputWidth; ++col) indices[row * kOutputWidth + col] = col; | ||
| 1167 | + } | ||
| 1168 | + std::fill_n(output, kRows * kOutputWidth, 0.0F); | ||
| 1169 | + AutofuseTilingData tiling_data{}; | ||
| 1170 | + uint32_t workspace_size = 0U; | ||
| 1171 | + uint32_t block_dim = 48U; | ||
| 1172 | + ASSERT_EQ(AutofuseTiling(&tiling_data, &workspace_size, &block_dim, 48U, 192U * 1024U), 0); | ||
| 1173 | + void *workspace = workspace_size == 0U ? nullptr : AscendC::GmAlloc(workspace_size); | ||
| 1174 | + ASSERT_TRUE(workspace_size == 0U || workspace != nullptr); | ||
| 1175 | + AscendC::SetKernelMode(KernelMode::AIV_MODE); | ||
| 1176 | + ICPU_RUN_KF(user_add_gather, block_dim, reinterpret_cast<uint8_t *>(input0), reinterpret_cast<uint8_t *>(input1), | ||
| 1177 | + reinterpret_cast<uint8_t *>(indices), reinterpret_cast<uint8_t *>(output), | ||
| 1178 | + reinterpret_cast<uint8_t *>(workspace), reinterpret_cast<uint8_t *>(&tiling_data)); | ||
| 1179 | + for (int32_t row = 0; row < kRows; ++row) { | ||
| 1180 | + for (int32_t col = 0; col < kOutputWidth; ++col) { | ||
| 1181 | + EXPECT_FLOAT_EQ(output[row * kOutputWidth + col], | ||
| 1182 | + input0[row * kInputWidth + col] + input1[row * kInputWidth + col]) | ||
| 1183 | + << "row=" << row << ", col=" << col; | ||
| 1184 | + } | ||
| 1185 | + } | ||
| 1186 | + if (workspace != nullptr) AscendC::GmFree(workspace); | ||
| 1187 | + AscendC::GmFree(input0); | ||
| 1188 | + AscendC::GmFree(input1); | ||
| 1189 | + AscendC::GmFree(indices); | ||
| 1190 | + AscendC::GmFree(output); | ||
| 1191 | +} | ||
| 1192 | + | ||
| 1193 | + | ||
| 844 | /** | 1194 | /** |
| 845 | * Copyright (c) 2026 Huawei Technologies Co., Ltd. | 1195 | * Copyright (c) 2026 Huawei Technologies Co., Ltd. |
| 846 | * This program is free software, you can redistribute it and/or modify it under the terms and conditions of | 1196 | * This program is free software, you can redistribute it and/or modify it under the terms and conditions of |
| @@ -861,7 +1211,15 @@ TEST(E2EIndirectLoadBroadcast, GeneratedKernelMatchesReference) { | |||
| 861 | extern "C" int64_t AutofuseTiling(AutofuseTilingData *, uint32_t *, uint32_t *, uint32_t, uint32_t); | 1211 | extern "C" int64_t AutofuseTiling(AutofuseTilingData *, uint32_t *, uint32_t *, uint32_t, uint32_t); |
| 862 | 1212 | ||
| 863 | 1213 | ||
| 864 | -#ifdef IL_EMBEDDING_REDUCE | 1214 | +#if defined(IL_GRAPH_HINT_SIMD_REPRO) |
| 1215 | +extern "C" __global__ __aicore__ void indirect_load_graph_hint_simd_repro(GM_ADDR input0, GM_ADDR input1, | ||
| 1216 | + GM_ADDR output, GM_ADDR workspace, | ||
| 1217 | + GM_ADDR tiling); | ||
| 1218 | + | ||
| 1219 | +extern "C" __global__ __aicore__ void indirect_load_graph_hint_reduce_simt_test(GM_ADDR input0, GM_ADDR input1, | ||
| 1220 | + GM_ADDR input2, GM_ADDR output, | ||
| 1221 | + GM_ADDR workspace, GM_ADDR tiling); | ||
| 1222 | + | ||
| 865 | extern "C" __global__ __aicore__ void indirect_load_embedding_reduce_simt_test(GM_ADDR input0, GM_ADDR input1, | 1223 | extern "C" __global__ __aicore__ void indirect_load_embedding_reduce_simt_test(GM_ADDR input0, GM_ADDR input1, |
| 866 | GM_ADDR input2, GM_ADDR output, | 1224 | GM_ADDR input2, GM_ADDR output, |
| 867 | GM_ADDR workspace, GM_ADDR tiling); | 1225 | GM_ADDR workspace, GM_ADDR tiling); |
| @@ -878,7 +1236,122 @@ extern "C" __global__ __aicore__ void indirect_load_add_il_reduce_test(GM_ADDR i | |||
| 878 | 1236 | ||
| 879 | namespace { | 1237 | namespace { |
| 880 | 1238 | ||
| 881 | -#ifdef IL_EMBEDDING_REDUCE | 1239 | +#if defined(IL_GRAPH_HINT_SIMD_REPRO) |
| 1240 | +constexpr int32_t kGraphHintSimdRows = 30; | ||
| 1241 | +constexpr int32_t kGraphHintSimdIndexColumns = 3; | ||
| 1242 | +constexpr int32_t kGraphHintSimdInner = 23; | ||
| 1243 | +constexpr int32_t kGraphHintSimdInputRows = 6; | ||
| 1244 | + | ||
| 1245 | +TEST(E2EIndirectLoadGraphHintSimdRepro, GeneratedKernelMatchesReference) { | ||
| 1246 | + constexpr int64_t input_count = | ||
| 1247 | + static_cast<int64_t>(kGraphHintSimdRows) * kGraphHintSimdInputRows * kGraphHintSimdInner; | ||
| 1248 | + constexpr int64_t index_count = kGraphHintSimdIndexColumns; | ||
| 1249 | + constexpr int64_t output_count = static_cast<int64_t>(kGraphHintSimdRows) * kGraphHintSimdIndexColumns; | ||
| 1250 | + auto *input = static_cast<float *>(AscendC::GmAlloc(sizeof(float) * input_count)); | ||
| 1251 | + auto *index = static_cast<int64_t *>(AscendC::GmAlloc(sizeof(int64_t) * index_count)); | ||
| 1252 | + auto *output = static_cast<float *>(AscendC::GmAlloc(sizeof(float) * output_count)); | ||
| 1253 | + ASSERT_NE(input, nullptr); | ||
| 1254 | + ASSERT_NE(index, nullptr); | ||
| 1255 | + ASSERT_NE(output, nullptr); | ||
| 1256 | + | ||
| 1257 | + for (int32_t row = 0; row < kGraphHintSimdRows; ++row) { | ||
| 1258 | + for (int32_t column = 0; column < kGraphHintSimdInputRows; ++column) { | ||
| 1259 | + for (int32_t inner = 0; inner < kGraphHintSimdInner; ++inner) { | ||
| 1260 | + input[(static_cast<int64_t>(row) * kGraphHintSimdInputRows + column) * kGraphHintSimdInner + inner] = | ||
| 1261 | + static_cast<float>((row * kGraphHintSimdInputRows + column) * kGraphHintSimdInner + inner); | ||
| 1262 | + } | ||
| 1263 | + } | ||
| 1264 | + } | ||
| 1265 | + std::fill_n(index, index_count, static_cast<int64_t>(5)); | ||
| 1266 | + std::fill_n(output, output_count, 0.0F); | ||
| 1267 | + | ||
| 1268 | + AutofuseTilingData tiling_data{}; | ||
| 1269 | + uint32_t workspace_size = 0U; | ||
| 1270 | + uint32_t block_dim = 48U; | ||
| 1271 | + ASSERT_EQ(AutofuseTiling(&tiling_data, &workspace_size, &block_dim, 48U, 192U * 1024U), 0); | ||
| 1272 | + void *workspace = workspace_size == 0U ? nullptr : AscendC::GmAlloc(workspace_size); | ||
| 1273 | + ASSERT_TRUE(workspace_size == 0U || workspace != nullptr); | ||
| 1274 | + | ||
| 1275 | + AscendC::SetKernelMode(KernelMode::AIV_MODE); | ||
| 1276 | + ICPU_RUN_KF(indirect_load_graph_hint_simd_repro, block_dim, reinterpret_cast<uint8_t *>(input), | ||
| 1277 | + reinterpret_cast<uint8_t *>(index), reinterpret_cast<uint8_t *>(output), | ||
| 1278 | + reinterpret_cast<uint8_t *>(workspace), reinterpret_cast<uint8_t *>(&tiling_data)); | ||
| 1279 | + | ||
| 1280 | + for (int32_t row = 0; row < kGraphHintSimdRows; ++row) { | ||
| 1281 | + for (int32_t column = 0; column < kGraphHintSimdIndexColumns; ++column) { | ||
| 1282 | + float expected = 0.0F; | ||
| 1283 | + for (int32_t inner = 0; inner < kGraphHintSimdInner; ++inner) { | ||
| 1284 | + expected += | ||
| 1285 | + input[(static_cast<int64_t>(row) * kGraphHintSimdInputRows + index[column]) * kGraphHintSimdInner + inner]; | ||
| 1286 | + } | ||
| 1287 | + EXPECT_FLOAT_EQ(output[row * kGraphHintSimdIndexColumns + column], expected) | ||
| 1288 | + << "row=" << row << ", column=" << column; | ||
| 1289 | + } | ||
| 1290 | + } | ||
| 1291 | + | ||
| 1292 | + if (workspace != nullptr) { | ||
| 1293 | + AscendC::GmFree(workspace); | ||
| 1294 | + } | ||
| 1295 | + AscendC::GmFree(input); | ||
| 1296 | + AscendC::GmFree(index); | ||
| 1297 | + AscendC::GmFree(output); | ||
| 1298 | +} | ||
| 1299 | + | ||
| 1300 | +constexpr int32_t kGraphHintRows = 8; | ||
| 1301 | +constexpr int32_t kGraphHintColumns = 50; | ||
| 1302 | +constexpr int32_t kGraphHintTableRows = 1353406; | ||
| 1303 | +constexpr int32_t kGraphHintTableStride = 8; | ||
| 1304 | + | ||
| 1305 | +TEST(E2EIndirectLoadGraphHintReduce, GeneratedKernelMatchesReference) { | ||
| 1306 | + auto *index0 = static_cast<int64_t *>(AscendC::GmAlloc(sizeof(int64_t) * kGraphHintColumns)); | ||
| 1307 | + auto *table = static_cast<float *>( | ||
| 1308 | + AscendC::GmAlloc(sizeof(float) * (static_cast<int64_t>(kGraphHintTableRows - 1) * kGraphHintTableStride + 1))); | ||
| 1309 | + auto *index2 = static_cast<int64_t *>(AscendC::GmAlloc(sizeof(int64_t) * kGraphHintColumns)); | ||
| 1310 | + auto *output = static_cast<float *>(AscendC::GmAlloc(sizeof(float) * kGraphHintRows)); | ||
| 1311 | + ASSERT_NE(index0, nullptr); | ||
| 1312 | + ASSERT_NE(table, nullptr); | ||
| 1313 | + ASSERT_NE(index2, nullptr); | ||
| 1314 | + ASSERT_NE(output, nullptr); | ||
| 1315 | + | ||
| 1316 | + for (int32_t row = 0; row < kGraphHintTableRows; ++row) { | ||
| 1317 | + table[static_cast<int64_t>(row) * kGraphHintTableStride] = static_cast<float>((row % 97) * 0.25F + 1.0F); | ||
| 1318 | + } | ||
| 1319 | + for (int32_t column = 0; column < kGraphHintColumns; ++column) { | ||
| 1320 | + index0[column] = (column % 3 == 0) ? -1 : static_cast<int64_t>(column * 10000 + 7); | ||
| 1321 | + index2[column] = static_cast<int64_t>(column * 20000 + 11); | ||
| 1322 | + } | ||
| 1323 | + | ||
| 1324 | + AutofuseTilingData tiling_data{}; | ||
| 1325 | + uint32_t workspace_size = 0; | ||
| 1326 | + uint32_t block_dim = 48; | ||
| 1327 | + ASSERT_EQ(AutofuseTiling(&tiling_data, &workspace_size, &block_dim, 48U, 192U * 1024U), 0); | ||
| 1328 | + void *workspace = workspace_size == 0U ? nullptr : AscendC::GmAlloc(workspace_size); | ||
| 1329 | + ASSERT_TRUE(workspace_size == 0U || workspace != nullptr); | ||
| 1330 | + | ||
| 1331 | + AscendC::SetKernelMode(KernelMode::AIV_MODE); | ||
| 1332 | + ICPU_RUN_KF(indirect_load_graph_hint_reduce_simt_test, block_dim, reinterpret_cast<uint8_t *>(index0), | ||
| 1333 | + reinterpret_cast<uint8_t *>(table), reinterpret_cast<uint8_t *>(index2), | ||
| 1334 | + reinterpret_cast<uint8_t *>(output), reinterpret_cast<uint8_t *>(workspace), | ||
| 1335 | + reinterpret_cast<uint8_t *>(&tiling_data)); | ||
| 1336 | + | ||
| 1337 | + for (int32_t row = 0; row < kGraphHintRows; ++row) { | ||
| 1338 | + float expected = 0.0F; | ||
| 1339 | + for (int32_t column = 0; column < kGraphHintColumns; ++column) { | ||
| 1340 | + const int64_t selected = index0[column] == -1 ? index2[column] : index0[column]; | ||
| 1341 | + expected += table[selected * kGraphHintTableStride]; | ||
| 1342 | + } | ||
| 1343 | + EXPECT_FLOAT_EQ(output[row], expected) << "row=" << row; | ||
| 1344 | + } | ||
| 1345 | + | ||
| 1346 | + if (workspace != nullptr) { | ||
| 1347 | + AscendC::GmFree(workspace); | ||
| 1348 | + } | ||
| 1349 | + AscendC::GmFree(index0); | ||
| 1350 | + AscendC::GmFree(table); | ||
| 1351 | + AscendC::GmFree(index2); | ||
| 1352 | + AscendC::GmFree(output); | ||
| 1353 | +} | ||
| 1354 | + | ||
| 882 | constexpr int32_t kEmbRows = 2; | 1355 | constexpr int32_t kEmbRows = 2; |
| 883 | constexpr int32_t kEmbColumns = 2; | 1356 | constexpr int32_t kEmbColumns = 2; |
| 884 | constexpr int32_t kEmbReduceSize = 2; | 1357 | constexpr int32_t kEmbReduceSize = 2; |
| @@ -18,6 +18,17 @@ | |||
| 18 | 18 | ||
| 19 | 19 | ||
| 20 | namespace AscendC { | 20 | namespace AscendC { |
| 21 | +template <int32_t Rank> | ||
| 22 | +struct IndirectLoadSimdStridedParams { | ||
| 23 | + uint32_t logical_size; | ||
| 24 | + uint32_t physical_size; | ||
| 25 | + int64_t output_offset; | ||
| 26 | + int64_t index_sizes[Rank]; | ||
| 27 | + int64_t input_strides[Rank]; | ||
| 28 | + int64_t index_strides[Rank]; | ||
| 29 | + int64_t output_strides[Rank]; | ||
| 30 | +}; | ||
| 31 | + | ||
| 21 | namespace Internal { | 32 | namespace Internal { |
| 22 | template <int32_t Dim, int32_t Axis, int32_t StrideBase> | 33 | template <int32_t Dim, int32_t Axis, int32_t StrideBase> |
| 23 | struct IndirectLoadSimdInnerOffset { | 34 | struct IndirectLoadSimdInnerOffset { |
| @@ -343,75 +354,158 @@ __aicore__ inline bool TryIndirectLoadSimdEmbedding(const LocalTensor<X> &x, con | |||
| 343 | return false; | 354 | return false; |
| 344 | } | 355 | } |
| 345 | 356 | ||
| 346 | -template <typename X, typename Index, int32_t Rank, int32_t Axis, typename... ShapeArgs> | 357 | +template <int32_t Rank, int32_t Axis> |
| 347 | -__aicore__ inline void IndirectLoadSimdStridedImpl(const LocalTensor<X> &x, const LocalTensor<Index> &index, | 358 | +struct IndirectLoadSimdStridedContext { |
| 348 | - const LocalTensor<X> &y, const LocalTensor<uint8_t> &tmp, | 359 | + const int64_t *shape; |
| 349 | - uint32_t actual_size, int64_t output_offset, | 360 | + const int64_t *output_strides; |
| 350 | - ShapeArgs... shape_args) { | 361 | + uint32_t logical_size; |
| 351 | - static_assert(Rank > 0 && Axis >= 0 && Axis < Rank, "IndirectLoad SIMD rank or axis is invalid."); | 362 | + uint32_t physical_size; |
| 352 | - static_assert(sizeof...(ShapeArgs) == static_cast<size_t>(3 * Rank), "IndirectLoad SIMD shape is invalid."); | 363 | + int64_t output_offset; |
| 353 | - const int64_t shape[] = {static_cast<int64_t>(shape_args)...}; | 364 | + int64_t index_inner; |
| 365 | + int64_t output_slice_count; | ||
| 366 | + int64_t input_window_base; | ||
| 367 | + int64_t index_window_base; | ||
| 368 | +}; | ||
| 369 | + | ||
| 370 | +template <int32_t Rank> | ||
| 371 | +__aicore__ inline void BuildIndirectLoadSimdStridedShape(int64_t (&shape)[3 * Rank], | ||
| 372 | + const IndirectLoadSimdStridedParams<Rank> ¶ms) { | ||
| 373 | + for (int32_t dim = 0; dim < Rank; ++dim) { | ||
| 374 | + shape[dim] = params.index_sizes[dim]; | ||
| 375 | + shape[Rank + dim] = params.input_strides[dim]; | ||
| 376 | + shape[2 * Rank + dim] = params.index_strides[dim]; | ||
| 377 | + } | ||
| 378 | +} | ||
| 379 | + | ||
| 380 | +template <int32_t Rank, int32_t Axis> | ||
| 381 | +__aicore__ inline IndirectLoadSimdStridedContext<Rank, Axis> MakeIndirectLoadSimdStridedContext( | ||
| 382 | + const int64_t *shape, const IndirectLoadSimdStridedParams<Rank> ¶ms) { | ||
| 354 | const int64_t index_inner = Internal::IndirectLoadSimdInnerSize<Rank - 1, Axis>::Call(shape); | 383 | const int64_t index_inner = Internal::IndirectLoadSimdInnerSize<Rank - 1, Axis>::Call(shape); |
| 355 | const int64_t output_slice_count = shape[Axis] * index_inner; | 384 | const int64_t output_slice_count = shape[Axis] * index_inner; |
| 356 | - const int64_t outer_begin = output_offset / output_slice_count; | 385 | + const int64_t outer_begin = params.output_offset / output_slice_count; |
| 357 | int64_t input_window_base = 0; | 386 | int64_t input_window_base = 0; |
| 358 | int64_t index_window_base = 0; | 387 | int64_t index_window_base = 0; |
| 359 | if constexpr (Axis > 0) { | 388 | if constexpr (Axis > 0) { |
| 360 | input_window_base = Internal::IndirectLoadSimdOuterOffset<Axis - 1, Rank>::Call(outer_begin, shape); | 389 | input_window_base = Internal::IndirectLoadSimdOuterOffset<Axis - 1, Rank>::Call(outer_begin, shape); |
| 361 | index_window_base = Internal::IndirectLoadSimdOuterOffset<Axis - 1, 2 * Rank>::Call(outer_begin, shape); | 390 | index_window_base = Internal::IndirectLoadSimdOuterOffset<Axis - 1, 2 * Rank>::Call(outer_begin, shape); |
| 362 | } | 391 | } |
| 363 | - if (TryIndirectLoadSimdEmbedding<X, Index, Rank, Axis>(x, index, y, actual_size, output_offset, shape)) { | 392 | + return {shape, params.output_strides, params.logical_size, params.physical_size, params.output_offset, |
| 364 | - return; | 393 | + index_inner, output_slice_count, input_window_base, index_window_base}; |
| 394 | +} | ||
| 395 | + | ||
| 396 | +template <int32_t Rank, int32_t Axis> | ||
| 397 | +__aicore__ inline int64_t GetIndirectLoadSimdStridedIndexOffset(int64_t global_idx, | ||
| 398 | + const IndirectLoadSimdStridedContext<Rank, Axis> &ctx) { | ||
| 399 | + const int64_t outer_global = global_idx / ctx.output_slice_count; | ||
| 400 | + const int64_t tail = global_idx % ctx.output_slice_count; | ||
| 401 | + const int64_t axis_coord = tail / ctx.index_inner; | ||
| 402 | + const int64_t inner = tail % ctx.index_inner; | ||
| 403 | + int64_t index_offset = axis_coord * ctx.shape[2 * Rank + Axis]; | ||
| 404 | + if constexpr (Axis > 0) { | ||
| 405 | + index_offset += Internal::IndirectLoadSimdOuterOffset<Axis - 1, 2 * Rank>::Call(outer_global, ctx.shape) - | ||
| 406 | + ctx.index_window_base; | ||
| 365 | } | 407 | } |
| 366 | - LocalTensor<uint32_t> offsets = tmp.template ReinterpretCast<uint32_t>(); | 408 | + if constexpr (Axis + 1 < Rank) { |
| 367 | - for (int64_t i = 0; i < actual_size; ++i) { | 409 | + index_offset += Internal::IndirectLoadSimdInnerOffset<Rank - 1, Axis, 2 * Rank>::Call(inner, ctx.shape); |
| 368 | - const int64_t global_idx = output_offset + i; | 410 | + } |
| 369 | - const int64_t outer_global = global_idx / output_slice_count; | 411 | + return index_offset; |
| 370 | - const int64_t tail = global_idx % output_slice_count; | 412 | +} |
| 371 | - const int64_t axis_coord = tail / index_inner; | 413 | + |
| 372 | - const int64_t inner = tail % index_inner; | 414 | +template <int32_t Rank, int32_t Axis> |
| 373 | - int64_t index_offset = axis_coord * shape[2 * Rank + Axis]; | 415 | +__aicore__ inline int64_t GetIndirectLoadSimdStridedInputOffset(int64_t global_idx, int64_t index_value, |
| 374 | - if constexpr (Axis > 0) { | 416 | + const IndirectLoadSimdStridedContext<Rank, Axis> &ctx) { |
| 375 | - index_offset += | 417 | + const int64_t outer_global = global_idx / ctx.output_slice_count; |
| 376 | - Internal::IndirectLoadSimdOuterOffset<Axis - 1, 2 * Rank>::Call(outer_global, shape) - index_window_base; | 418 | + const int64_t tail = global_idx % ctx.output_slice_count; |
| 377 | - } | 419 | + const int64_t inner = tail % ctx.index_inner; |
| 378 | - if constexpr (Axis + 1 < Rank) { | 420 | + int64_t input_inner_offset = 0; |
| 379 | - index_offset += Internal::IndirectLoadSimdInnerOffset<Rank - 1, Axis, 2 * Rank>::Call(inner, shape); | 421 | + int64_t input_outer_offset = 0; |
| 380 | - } | 422 | + if constexpr (Axis + 1 < Rank) { |
| 423 | + input_inner_offset = Internal::IndirectLoadSimdInnerOffset<Rank - 1, Axis, Rank>::Call(inner, ctx.shape); | ||
| 424 | + } | ||
| 425 | + if constexpr (Axis > 0) { | ||
| 426 | + input_outer_offset = | ||
| 427 | + Internal::IndirectLoadSimdOuterOffset<Axis - 1, Rank>::Call(outer_global, ctx.shape) - ctx.input_window_base; | ||
| 428 | + } | ||
| 429 | + return input_outer_offset + index_value * ctx.shape[Rank + Axis] + input_inner_offset; | ||
| 430 | +} | ||
| 431 | + | ||
| 432 | +template <typename X, typename Index, int32_t Rank, int32_t Axis> | ||
| 433 | +__aicore__ inline void BuildIndirectLoadSimdStridedOffsets(const LocalTensor<Index> &index, | ||
| 434 | + const LocalTensor<uint32_t> &offsets, | ||
| 435 | + const IndirectLoadSimdStridedContext<Rank, Axis> &ctx) { | ||
| 436 | + for (uint32_t i = 0; i < ctx.logical_size; ++i) { | ||
| 437 | + const int64_t global_idx = ctx.output_offset + static_cast<int64_t>(i); | ||
| 438 | + const int64_t index_offset = GetIndirectLoadSimdStridedIndexOffset<Rank, Axis>(global_idx, ctx); | ||
| 381 | const int64_t index_value = static_cast<int64_t>(index.GetValue(index_offset)); | 439 | const int64_t index_value = static_cast<int64_t>(index.GetValue(index_offset)); |
| 382 | - int64_t input_inner_offset = 0; | 440 | + const int64_t source_offset = GetIndirectLoadSimdStridedInputOffset<Rank, Axis>(global_idx, index_value, ctx); |
| 383 | - if constexpr (Axis + 1 < Rank) { | 441 | + offsets.SetValue(i, static_cast<uint32_t>(source_offset * sizeof(X))); |
| 384 | - input_inner_offset = Internal::IndirectLoadSimdInnerOffset<Rank - 1, Axis, Rank>::Call(inner, shape); | ||
| 385 | - } | ||
| 386 | - int64_t input_outer_offset = 0; | ||
| 387 | - if constexpr (Axis > 0) { | ||
| 388 | - input_outer_offset = | ||
| 389 | - Internal::IndirectLoadSimdOuterOffset<Axis - 1, Rank>::Call(outer_global, shape) - input_window_base; | ||
| 390 | - } | ||
| 391 | - const int64_t src_idx = input_outer_offset + index_value * shape[Rank + Axis] + input_inner_offset; | ||
| 392 | - offsets.SetValue(i, static_cast<uint32_t>(src_idx * sizeof(X))); | ||
| 393 | } | 442 | } |
| 443 | +} | ||
| 444 | + | ||
| 445 | +template <int32_t Rank, int32_t Axis> | ||
| 446 | +__aicore__ inline int64_t GetIndirectLoadSimdStridedOutputOffset( | ||
| 447 | + int64_t global_idx, const IndirectLoadSimdStridedContext<Rank, Axis> &ctx) { | ||
| 448 | + int64_t current = global_idx; | ||
| 449 | + int64_t base = ctx.output_offset; | ||
| 450 | + int64_t output_offset = 0; | ||
| 451 | + for (int32_t dim = Rank - 1; dim >= 0; --dim) { | ||
| 452 | + const int64_t coord = current % ctx.shape[dim]; | ||
| 453 | + const int64_t base_coord = base % ctx.shape[dim]; | ||
| 454 | + current /= ctx.shape[dim]; | ||
| 455 | + base /= ctx.shape[dim]; | ||
| 456 | + output_offset += (coord - base_coord) * ctx.output_strides[dim]; | ||
| 457 | + } | ||
| 458 | + return output_offset; | ||
| 459 | +} | ||
| 460 | + | ||
| 461 | +template <typename X, int32_t Rank, int32_t Axis> | ||
| 462 | +__aicore__ inline void ScatterIndirectLoadSimdStridedOutput(const LocalTensor<X> &y, | ||
| 463 | + const IndirectLoadSimdStridedContext<Rank, Axis> &ctx) { | ||
| 464 | + // Gather writes packed logical values first. Repack from the end so padding lanes do not | ||
| 465 | + // overwrite a source value that is still needed by the supported compact output layouts. | ||
| 466 | + for (int64_t i = static_cast<int64_t>(ctx.logical_size) - 1; i >= 0; --i) { | ||
| 467 | + const int64_t output_offset = GetIndirectLoadSimdStridedOutputOffset<Rank, Axis>(ctx.output_offset + i, ctx); | ||
| 468 | + y.SetValue(output_offset, y.GetValue(i)); | ||
| 469 | + } | ||
| 470 | +} | ||
| 471 | + | ||
| 472 | +// Strided SIMD output may contain alignment holes (for example, a logical 23-element row | ||
| 473 | +// is stored with a physical stride of 24). Keep the gather packed and expand it into the | ||
| 474 | +// padded destination layout afterwards. | ||
| 475 | +template <typename X, typename Index, int32_t Rank, int32_t Axis> | ||
| 476 | +__aicore__ inline void IndirectLoadSimdStridedImpl(const LocalTensor<X> &x, const LocalTensor<Index> &index, | ||
| 477 | + const LocalTensor<X> &y, const LocalTensor<uint8_t> &tmp, | ||
| 478 | + const IndirectLoadSimdStridedParams<Rank> ¶ms) { | ||
| 479 | + static_assert(Rank > 0 && Axis >= 0 && Axis < Rank, "IndirectLoad SIMD rank or axis is invalid."); | ||
| 480 | + int64_t shape[3 * Rank]; | ||
| 481 | + BuildIndirectLoadSimdStridedShape(shape, params); | ||
| 482 | + const auto context = MakeIndirectLoadSimdStridedContext<Rank, Axis>(shape, params); | ||
| 483 | + LocalTensor<uint32_t> offsets = tmp.template ReinterpretCast<uint32_t>(); | ||
| 484 | + BuildIndirectLoadSimdStridedOffsets<X, Index, Rank, Axis>(index, offsets, context); | ||
| 394 | int32_t offset_event_id = static_cast<int32_t>(GetTPipePtr()->FetchEventID(AscendC::HardEvent::S_V)); | 485 | int32_t offset_event_id = static_cast<int32_t>(GetTPipePtr()->FetchEventID(AscendC::HardEvent::S_V)); |
| 395 | AscendC::SetFlag<AscendC::HardEvent::S_V>(offset_event_id); | 486 | AscendC::SetFlag<AscendC::HardEvent::S_V>(offset_event_id); |
| 396 | AscendC::WaitFlag<AscendC::HardEvent::S_V>(offset_event_id); | 487 | AscendC::WaitFlag<AscendC::HardEvent::S_V>(offset_event_id); |
| 397 | - Gather(y, x, offsets, static_cast<uint32_t>(0), actual_size); | 488 | + Gather(y, x, offsets, static_cast<uint32_t>(0), context.logical_size); |
| 489 | + ScatterIndirectLoadSimdStridedOutput<X, Rank, Axis>(y, context); | ||
| 398 | } | 490 | } |
| 399 | } // namespace Internal | 491 | } // namespace Internal |
| 400 | 492 | ||
| 401 | -template <typename X, typename Index, int32_t Rank, int32_t Axis, typename FirstArg, typename... Args> | 493 | +template <typename X, typename Index, int32_t Rank, int32_t Axis, typename... ShapeArgs> |
| 402 | __aicore__ inline void IndirectLoadSimd(const LocalTensor<X> &x, const LocalTensor<Index> &index, | 494 | __aicore__ inline void IndirectLoadSimd(const LocalTensor<X> &x, const LocalTensor<Index> &index, |
| 403 | - const LocalTensor<X> &y, FirstArg first_arg, Args... args) { | 495 | + const LocalTensor<X> &y, uint32_t actual_size, int64_t output_offset, |
| 496 | + uint32_t input_actual_size, int64_t input_axis, ShapeArgs... shape_args) { | ||
| 404 | static_assert(Rank > 0 && Axis >= 0 && Axis < Rank, "IndirectLoad SIMD rank or axis is invalid."); | 497 | static_assert(Rank > 0 && Axis >= 0 && Axis < Rank, "IndirectLoad SIMD rank or axis is invalid."); |
| 405 | - constexpr bool has_tmp = std::is_same_v<std::decay_t<FirstArg>, LocalTensor<uint8_t>>; | 498 | + static_assert(sizeof...(ShapeArgs) == static_cast<size_t>(2 * Rank), "IndirectLoad SIMD shape is invalid."); |
| 406 | - if constexpr (has_tmp) { | 499 | + Internal::IndirectLoadSimdDenseImpl<X, Index, Rank, Axis>(x, index, y, actual_size, output_offset, input_actual_size, |
| 407 | - static_assert(sizeof...(Args) == static_cast<size_t>(2 + 3 * Rank), | 500 | + input_axis, shape_args...); |
| 408 | - "IndirectLoad SIMD strided arguments are invalid."); | 501 | +} |
| 409 | - Internal::IndirectLoadSimdStridedImpl<X, Index, Rank, Axis>(x, index, y, first_arg, args...); | 502 | + |
| 410 | - } else { | 503 | +template <typename X, typename Index, int32_t Rank, int32_t Axis> |
| 411 | - static_assert(sizeof...(Args) == static_cast<size_t>(3 + 2 * Rank), | 504 | +__aicore__ inline void IndirectLoadSimdStrided(const LocalTensor<X> &x, const LocalTensor<Index> &index, |
| 412 | - "IndirectLoad SIMD dense arguments are invalid."); | 505 | + const LocalTensor<X> &y, const LocalTensor<uint8_t> &tmp, |
| 413 | - Internal::IndirectLoadSimdDenseImpl<X, Index, Rank, Axis>(x, index, y, first_arg, args...); | 506 | + const IndirectLoadSimdStridedParams<Rank> ¶ms) { |
| 414 | - } | 507 | + static_assert(Rank > 0 && Axis >= 0 && Axis < Rank, "IndirectLoad SIMD rank or axis is invalid."); |
| 508 | + Internal::IndirectLoadSimdStridedImpl<X, Index, Rank, Axis>(x, index, y, tmp, params); | ||
| 415 | } | 509 | } |
| 416 | 510 | ||
| 417 | template <typename X, typename Index, int32_t Rank, int32_t Axis, typename... ShapeArgs> | 511 | template <typename X, typename Index, int32_t Rank, int32_t Axis, typename... ShapeArgs> |
| @@ -3659,8 +3659,8 @@ class IndirectLoadAscIrCodegenImplV2 : public AscIrCodegenV2 { | |||
| 3659 | return false; | 3659 | return false; |
| 3660 | } | 3660 | } |
| 3661 | [[nodiscard]] std::vector<std::string> LoadApiHeaderFiles([[maybe_unused]] bool is_dynamic) const override { | 3661 | [[nodiscard]] std::vector<std::string> LoadApiHeaderFiles([[maybe_unused]] bool is_dynamic) const override { |
| 3662 | - return {"indirect_load_simd_policy_reg_base.h", "indirect_load_simd_reg_base.h", "indirect_load_sk_reg_base.h", | 3662 | + return {"datacopy_reg_base.h", "indirect_load_simd_policy_reg_base.h", "indirect_load_simd_reg_base.h", |
| 3663 | - "indirect_load_simt_reg_base.h"}; | 3663 | + "indirect_load_sk_reg_base.h", "indirect_load_simt_reg_base.h"}; |
| 3664 | } | 3664 | } |
| 3665 | [[nodiscard]] std::vector<std::string> IncludeApiHeaderFiles() const override { | 3665 | [[nodiscard]] std::vector<std::string> IncludeApiHeaderFiles() const override { |
| 3666 | return {"basic_api/kernel_operator_vec_gather_intf.h", "basic_api/reg_compute/kernel_reg_compute_intf.h", | 3666 | return {"basic_api/kernel_operator_vec_gather_intf.h", "basic_api/reg_compute/kernel_reg_compute_intf.h", |
| @@ -69,6 +69,22 @@ struct SimtCodegenPlan { | |||
| 69 | uint64_t index_stride_mask = 0U; | 69 | uint64_t index_stride_mask = 0U; |
| 70 | }; | 70 | }; |
| 71 | 71 | ||
| 72 | +Status GenerateSimtContextDefinition(const std::string &context_name, const std::vector<SimtGmTensor> &gm_tensors, | ||
| 73 | + std::stringstream &ss) { | ||
| 74 | + ss << "struct " << context_name << " {" << std::endl; | ||
| 75 | + for (const SimtGmTensor &tensor : gm_tensors) { | ||
| 76 | + std::string dtype; | ||
| 77 | + GE_ASSERT_SUCCESS(Tensor::DtypeName(tensor.dtype, dtype)); | ||
| 78 | + if (tensor.is_scalar) { | ||
| 79 | + ss << " " << dtype << " " << kSimtValueNamePrefix << tensor.value_tensor_id << ";" << std::endl; | ||
| 80 | + } else { | ||
| 81 | + ss << " __gm__ " << dtype << " *" << kSimtGmFieldNamePrefix << tensor.value_tensor_id << ";" << std::endl; | ||
| 82 | + } | ||
| 83 | + } | ||
| 84 | + ss << "};" << std::endl; | ||
| 85 | + return af::SUCCESS; | ||
| 86 | +} | ||
| 87 | + | ||
| 72 | af::Status EmitSimtScalarExpr(const ascir::NodeView &node, const std::vector<std::string> &inputs, std::string &expr) { | 88 | af::Status EmitSimtScalarExpr(const ascir::NodeView &node, const std::vector<std::string> &inputs, std::string &expr) { |
| 73 | GE_ASSERT_NOTNULL(node, "SIMT scalar node is null."); | 89 | GE_ASSERT_NOTNULL(node, "SIMT scalar node is null."); |
| 74 | GE_ASSERT_TRUE(!inputs.empty() && inputs.size() == node->inputs.Size(), | 90 | GE_ASSERT_TRUE(!inputs.empty() && inputs.size() == node->inputs.Size(), |
| @@ -436,7 +452,8 @@ af::Status CollectSimtBackwardNodes(const af::AscNodePtr &root, const ascir::Nod | |||
| 436 | if (current == nullptr || current == indirect_load || !nodes.emplace(current.get()).second) { | 452 | if (current == nullptr || current == indirect_load || !nodes.emplace(current.get()).second) { |
| 437 | continue; | 453 | continue; |
| 438 | } | 454 | } |
| 439 | - if (af::ops::IsOps<af::ascir_op::Load>(current) || af::ops::IsOps<af::ascir_op::Scalar>(current)) { | 455 | + if (af::ops::IsOps<af::ascir_op::Load>(current) || af::ops::IsOps<af::ascir_op::Scalar>(current) || |
| 456 | + af::ops::IsOps<af::ascir_op::ScalarData>(current)) { | ||
| 440 | continue; | 457 | continue; |
| 441 | } | 458 | } |
| 442 | GE_ASSERT_TRUE(current->inputs.Size() > 0UL, "IndirectLoad SIMT node[%s] has no input.", current->GetNamePtr()); | 459 | GE_ASSERT_TRUE(current->inputs.Size() > 0UL, "IndirectLoad SIMT node[%s] has no input.", current->GetNamePtr()); |
| @@ -450,18 +467,51 @@ af::Status CollectSimtBackwardNodes(const af::AscNodePtr &root, const ascir::Nod | |||
| 450 | } | 467 | } |
| 451 | 468 | ||
| 452 | af::AscNodePtr FindSimtOutputStore(const ascir::NodeView &indirect_load) { | 469 | af::AscNodePtr FindSimtOutputStore(const ascir::NodeView &indirect_load) { |
| 453 | - for (af::AscNodePtr current = ascgen_utils::indirect_load::GetOnlyOutputConsumer(indirect_load); current != nullptr; | 470 | + // The SIMT evaluator can process a branched scalar region as long as all |
| 454 | - current = ascgen_utils::indirect_load::GetOnlyOutputConsumer(current)) { | 471 | + // branches merge into one final Store. Do not use GetOnlyOutputConsumer |
| 455 | - if (af::ops::IsOps<af::ascir_op::Store>(current)) { | 472 | + // here: it intentionally returns null at a fan-out and would reject this |
| 456 | - return current; | 473 | + // otherwise valid topology before backward-region collection runs. |
| 474 | + std::vector<af::AscNodePtr> pending; | ||
| 475 | + SimtNodeSet visited; | ||
| 476 | + for (const auto &out_node : indirect_load->GetOutDataNodes()) { | ||
| 477 | + const auto consumer = std::dynamic_pointer_cast<af::AscNode>(out_node); | ||
| 478 | + if (consumer != nullptr && visited.emplace(consumer.get()).second) { | ||
| 479 | + pending.emplace_back(consumer); | ||
| 457 | } | 480 | } |
| 458 | } | 481 | } |
| 459 | - return nullptr; | 482 | + af::AscNodePtr store; |
| 483 | + af::AscNodePtr fallback_store; | ||
| 484 | + for (size_t cursor = 0UL; cursor < pending.size(); ++cursor) { | ||
| 485 | + const auto ¤t = pending[cursor]; | ||
| 486 | + if (af::ops::IsOps<af::ascir_op::Store>(current)) { | ||
| 487 | + // Prefer the store whose producer belongs to the annotated SIMT main | ||
| 488 | + // chain. A fan-out branch may be materialized as a VectorFunc with the | ||
| 489 | + // auxiliary fanout role; selecting it would make the SIMT evaluator | ||
| 490 | + // parse an ordinary scheduled branch as scalar output code. | ||
| 491 | + if (fallback_store == nullptr) { | ||
| 492 | + fallback_store = current; | ||
| 493 | + } | ||
| 494 | + const auto producer = ascgen_utils::indirect_load::GetInputProducer(current, 0UL); | ||
| 495 | + if (store == nullptr && producer != nullptr && | ||
| 496 | + ascgen_utils::indirect_load::GetTemplateRole(producer) == | ||
| 497 | + ascgen_utils::indirect_load::TemplateRole::kSimtInlineTransform) { | ||
| 498 | + store = current; | ||
| 499 | + } | ||
| 500 | + continue; | ||
| 501 | + } | ||
| 502 | + for (const auto &out_node : current->GetOutDataNodes()) { | ||
| 503 | + const auto consumer = std::dynamic_pointer_cast<af::AscNode>(out_node); | ||
| 504 | + if (consumer != nullptr && visited.emplace(consumer.get()).second) { | ||
| 505 | + pending.emplace_back(consumer); | ||
| 506 | + } | ||
| 507 | + } | ||
| 508 | + } | ||
| 509 | + return store != nullptr ? store : fallback_store; | ||
| 460 | } | 510 | } |
| 461 | 511 | ||
| 462 | af::Status ValidateSimtRegionNode(const af::AscNodePtr &node) { | 512 | af::Status ValidateSimtRegionNode(const af::AscNodePtr &node) { |
| 463 | if (af::ops::IsOps<af::ascir_op::Load>(node) || af::ops::IsOps<af::ascir_op::Scalar>(node) || | 513 | if (af::ops::IsOps<af::ascir_op::Load>(node) || af::ops::IsOps<af::ascir_op::Scalar>(node) || |
| 464 | - af::ops::IsOps<af::ascir_op::Store>(node)) { | 514 | + af::ops::IsOps<af::ascir_op::ScalarData>(node) || af::ops::IsOps<af::ascir_op::Store>(node)) { |
| 465 | return af::SUCCESS; | 515 | return af::SUCCESS; |
| 466 | } | 516 | } |
| 467 | GE_ASSERT_TRUE(ascgen_utils::indirect_load::GetTemplateRole(node) == | 517 | GE_ASSERT_TRUE(ascgen_utils::indirect_load::GetTemplateRole(node) == |
| @@ -482,6 +532,16 @@ void AppendSimtGmTensor(const af::AscNodePtr &node, std::vector<SimtGmTensor> &g | |||
| 482 | } | 532 | } |
| 483 | } | 533 | } |
| 484 | 534 | ||
| 535 | +void AppendSimtScalarData(const af::AscNodePtr &node, std::vector<SimtGmTensor> &gm_tensors) { | ||
| 536 | + const auto output = node->outputs()[0]; | ||
| 537 | + const auto found = std::find_if(gm_tensors.begin(), gm_tensors.end(), [output](const SimtGmTensor &tensor) { | ||
| 538 | + return tensor.value_tensor_id == output->attr.mem.tensor_id; | ||
| 539 | + }); | ||
| 540 | + if (found == gm_tensors.end()) { | ||
| 541 | + gm_tensors.push_back({output->attr.mem.tensor_id, output->attr.mem.tensor_id, output->attr.dtype, true}); | ||
| 542 | + } | ||
| 543 | +} | ||
| 544 | + | ||
| 485 | af::Status CollectSimtRegionMetadata(const ascir::NodeView &indirect_load, std::vector<af::AscNodePtr> &index_nodes, | 545 | af::Status CollectSimtRegionMetadata(const ascir::NodeView &indirect_load, std::vector<af::AscNodePtr> &index_nodes, |
| 486 | std::vector<af::AscNodePtr> &output_nodes, std::vector<SimtGmTensor> &gm_tensors, | 546 | std::vector<af::AscNodePtr> &output_nodes, std::vector<SimtGmTensor> &gm_tensors, |
| 487 | af::AscNodePtr &store) { | 547 | af::AscNodePtr &store) { |
| @@ -511,6 +571,10 @@ af::Status CollectSimtRegionMetadata(const ascir::NodeView &indirect_load, std:: | |||
| 511 | (index_set.count(node.get()) != 0UL || output_set.count(node.get()) != 0UL)) { | 571 | (index_set.count(node.get()) != 0UL || output_set.count(node.get()) != 0UL)) { |
| 512 | AppendSimtGmTensor(node, gm_tensors); | 572 | AppendSimtGmTensor(node, gm_tensors); |
| 513 | } | 573 | } |
| 574 | + if (af::ops::IsOps<af::ascir_op::ScalarData>(node) && | ||
| 575 | + (index_set.count(node.get()) != 0UL || output_set.count(node.get()) != 0UL)) { | ||
| 576 | + AppendSimtScalarData(node, gm_tensors); | ||
| 577 | + } | ||
| 514 | } | 578 | } |
| 515 | return af::SUCCESS; | 579 | return af::SUCCESS; |
| 516 | } | 580 | } |
| @@ -532,6 +596,9 @@ af::Status CollectSimtMetadata(const ascir::NodeView &indirect_load, const af::A | |||
| 532 | if (af::ops::IsOps<af::ascir_op::Load>(node)) { | 596 | if (af::ops::IsOps<af::ascir_op::Load>(node)) { |
| 533 | AppendSimtGmTensor(node, gm_tensors); | 597 | AppendSimtGmTensor(node, gm_tensors); |
| 534 | } | 598 | } |
| 599 | + if (af::ops::IsOps<af::ascir_op::ScalarData>(node)) { | ||
| 600 | + AppendSimtScalarData(node, gm_tensors); | ||
| 601 | + } | ||
| 535 | } | 602 | } |
| 536 | return af::SUCCESS; | 603 | return af::SUCCESS; |
| 537 | } | 604 | } |
| @@ -588,6 +655,13 @@ af::Status GenerateSimtEvaluatorBody(const std::vector<af::AscNodePtr> &nodes, | |||
| 588 | GE_ASSERT_SUCCESS(EmitSimtScalarInput(node, values, ss)); | 655 | GE_ASSERT_SUCCESS(EmitSimtScalarInput(node, values, ss)); |
| 589 | continue; | 656 | continue; |
| 590 | } | 657 | } |
| 658 | + if (af::ops::IsOps<af::ascir_op::ScalarData>(node)) { | ||
| 659 | + const auto output = node->outputs()[0]; | ||
| 660 | + const std::string variable = | ||
| 661 | + "context." + std::string(kSimtValueNamePrefix) + std::to_string(output->attr.mem.tensor_id); | ||
| 662 | + values[output->attr.mem.tensor_id] = variable; | ||
| 663 | + continue; | ||
| 664 | + } | ||
| 591 | if (af::ops::IsOps<af::ascir_op::Store>(node)) { | 665 | if (af::ops::IsOps<af::ascir_op::Store>(node)) { |
| 592 | continue; | 666 | continue; |
| 593 | } | 667 | } |
| @@ -628,15 +702,139 @@ af::Status CalcVectorizedElementCount(const Tensor &tensor, af::Expression &elem | |||
| 628 | return af::SUCCESS; | 702 | return af::SUCCESS; |
| 629 | } | 703 | } |
| 630 | 704 | ||
| 705 | +std::string GetSimdLogicalOutputSize(const TPipe &tpipe, const Tensor &tensor) { | ||
| 706 | + std::stringstream ss; | ||
| 707 | + ss << "1"; | ||
| 708 | + for (size_t i = 0; i < tensor.vectorized_axis.size(); ++i) { | ||
| 709 | + const uint32_t axis_pos = tensor.vectorized_axis_pos[i]; | ||
| 710 | + if (axis_pos >= tensor.axis_size.size()) { | ||
| 711 | + continue; | ||
| 712 | + } | ||
| 713 | + const auto &axis = tpipe.tiler.GetAxis(tensor.vectorized_axis[i]); | ||
| 714 | + const auto &axis_size = tensor.axis_size[axis_pos]; | ||
| 715 | + const bool use_actual = axis.type == Axis::Type::kAxisTypeTileInner || | ||
| 716 | + af::SymbolicUtils::StaticCheckEq(axis_size, axis.size_expr) == af::TriBool::kTrue; | ||
| 717 | + ss << " * " << (use_actual ? "(" + axis.actual_size.Str() + ")" : "(" + tpipe.tiler.Size(axis_size) + ")"); | ||
| 718 | + } | ||
| 719 | + return ss.str(); | ||
| 720 | +} | ||
| 721 | + | ||
| 722 | +std::vector<ascir::SizeExpr> GetSimdOutputStrides(const Tensor &tensor, | ||
| 723 | + const ascgen_utils::indirect_load::LogicalTensorView &layout) { | ||
| 724 | + std::vector<ascir::SizeExpr> strides(layout.axis_ids.size(), af::sym::kSymbolOne); | ||
| 725 | + for (size_t dim = 0; dim < layout.axis_ids.size(); ++dim) { | ||
| 726 | + const auto it = std::find(tensor.vectorized_axis.begin(), tensor.vectorized_axis.end(), layout.axis_ids[dim]); | ||
| 727 | + if (it != tensor.vectorized_axis.end()) { | ||
| 728 | + strides[dim] = tensor.vectorized_strides[static_cast<size_t>(std::distance(tensor.vectorized_axis.begin(), it))]; | ||
| 729 | + continue; | ||
| 730 | + } | ||
| 731 | + const auto axis_it = std::find(tensor.axis.begin(), tensor.axis.end(), layout.axis_ids[dim]); | ||
| 732 | + if (axis_it != tensor.axis.end()) { | ||
| 733 | + strides[dim] = tensor.axis_strides[static_cast<size_t>(std::distance(tensor.axis.begin(), axis_it))]; | ||
| 734 | + } | ||
| 735 | + } | ||
| 736 | + return strides; | ||
| 737 | +} | ||
| 738 | + | ||
| 739 | +enum class SimdApiKind { | ||
| 740 | + kDense, | ||
| 741 | + kGather, | ||
| 742 | + kStrided, | ||
| 743 | +}; | ||
| 744 | + | ||
| 745 | +SimdApiKind SelectSimdApi(const ascgen_utils::indirect_load::TemplateLogicalView &logical_view, | ||
| 746 | + ascgen_utils::indirect_load::Implementation implementation) { | ||
| 747 | + const bool strided = logical_view.input.kind != ascgen_utils::indirect_load::IndirectLoadLayoutKind::kDense || | ||
| 748 | + logical_view.index.kind != ascgen_utils::indirect_load::IndirectLoadLayoutKind::kDense; | ||
| 749 | + if (strided) { | ||
| 750 | + return SimdApiKind::kStrided; | ||
| 751 | + } | ||
| 752 | + return implementation == ascgen_utils::indirect_load::Implementation::kGatherApi ? SimdApiKind::kGather | ||
| 753 | + : SimdApiKind::kDense; | ||
| 754 | +} | ||
| 755 | + | ||
| 756 | +const char *GetSimdApiName(SimdApiKind api_kind) { | ||
| 757 | + switch (api_kind) { | ||
| 758 | + case SimdApiKind::kStrided: | ||
| 759 | + return "IndirectLoadSimdStrided"; | ||
| 760 | + case SimdApiKind::kGather: | ||
| 761 | + return "IndirectLoadSimdGatherApi"; | ||
| 762 | + case SimdApiKind::kDense: | ||
| 763 | + return "IndirectLoadSimd"; | ||
| 764 | + } | ||
| 765 | + return "IndirectLoadSimd"; | ||
| 766 | +} | ||
| 767 | + | ||
| 768 | +void EmitSimdStridedParams(const TPipe &tpipe, const std::vector<ascir::AxisId> ¤t_axis, const Tensor &output, | ||
| 769 | + const LogicalTensorInfo &input_info, const LogicalTensorInfo &index_info, | ||
| 770 | + const ascgen_utils::indirect_load::LogicalTensorView &output_layout, size_t rank, | ||
| 771 | + std::stringstream &ss) { | ||
| 772 | + const std::string logical_output_size = GetSimdLogicalOutputSize(tpipe, output); | ||
| 773 | + const auto output_strides = GetSimdOutputStrides(output, output_layout); | ||
| 774 | + ss << " AscendC::IndirectLoadSimdStridedParams<" << rank << "> indirect_load_simd_params{static_cast<uint32_t>(" | ||
| 775 | + << logical_output_size << "), static_cast<uint32_t>(" << output.actual_size << "), " | ||
| 776 | + << tpipe.tiler.Offset(current_axis, output.axis, output.axis_strides) << ", {" | ||
| 777 | + << JoinSizeExprs(index_info.sizes, tpipe) << "}, {" << JoinSizeExprs(input_info.strides, tpipe) << "}, {" | ||
| 778 | + << JoinSizeExprs(index_info.strides, tpipe) << "}, {" << JoinSizeExprs(output_strides, tpipe) << "}};" | ||
| 779 | + << std::endl; | ||
| 780 | +} | ||
| 781 | + | ||
| 782 | +void EmitSimdInvocation(const TPipe &tpipe, const std::vector<ascir::AxisId> ¤t_axis, const Tensor &input, | ||
| 783 | + const Tensor &index, const Tensor &output, const LogicalTensorInfo &input_info, | ||
| 784 | + const LogicalTensorInfo &index_info, int64_t axis, size_t axis_pos, SimdApiKind api_kind, | ||
| 785 | + const std::string &input_dtype, const std::string &index_dtype, const std::string &tmp_name, | ||
| 786 | + std::stringstream &ss) { | ||
| 787 | + ss << " AscendC::" << GetSimdApiName(api_kind) << "<" << input_dtype << ", " << index_dtype << ", " | ||
| 788 | + << input_info.sizes.size() << ", " << axis << ">(\n"; | ||
| 789 | + ss << " " << input << ", " << index << ", " << output << ", "; | ||
| 790 | + if (api_kind == SimdApiKind::kStrided) { | ||
| 791 | + ss << tmp_name << ", indirect_load_simd_params"; | ||
| 792 | + } else { | ||
| 793 | + ss << output.actual_size << ", " << tpipe.tiler.Offset(current_axis, output.axis, output.axis_strides) << ", " | ||
| 794 | + << input.actual_size << ", " << tpipe.tiler.Size(input_info.sizes[axis_pos]) << ", " | ||
| 795 | + << JoinSizeExprs(index_info.sizes, tpipe) << ", " << JoinSizeExprs(input_info.strides, tpipe); | ||
| 796 | + } | ||
| 797 | + ss << ");" << std::endl; | ||
| 798 | +} | ||
| 799 | + | ||
| 800 | +void EmitSimtInvocation(const TPipe &tpipe, const std::string &input_dtype, const std::string &output_dtype, | ||
| 801 | + const std::string &outer_tb_var, const std::string &body_name, | ||
| 802 | + const LogicalTensorInfo &input_info, const LogicalTensorInfo &index_info, | ||
| 803 | + const SimtCodegenPlan &plan, int64_t axis, bool has_post_reduce, const Tensor *output_tensor, | ||
| 804 | + const af::Expression &output_element_count, std::stringstream &ss) { | ||
| 805 | + ss << " AscendC::IndirectLoadSimt<" << input_dtype << ", " << output_dtype << ", " << body_name << ", "; | ||
| 806 | + ss << GetSimtPolicyType(plan, input_info.sizes.size(), axis) << ">(" << std::endl; | ||
| 807 | + if (has_post_reduce) { | ||
| 808 | + const std::string output_elements = tpipe.tiler.Size(output_element_count); | ||
| 809 | + ss << " input_ptr, " << *output_tensor << ", context, static_cast<uint32_t>(" << output_elements << "), "; | ||
| 810 | + ss << "(static_cast<" << plan.offset_type << ">(block_dim_offset) + static_cast<" << plan.offset_type << ">(" | ||
| 811 | + << outer_tb_var << ")) * " << PromoteSizeExpr(output_elements, plan.offset_type); | ||
| 812 | + } else { | ||
| 813 | + ss << " input_ptr, y_ptr, context, static_cast<uint32_t>(" << outer_tb_var << "_loop_size), "; | ||
| 814 | + ss << "static_cast<" << plan.offset_type << ">(block_dim_offset)"; | ||
| 815 | + } | ||
| 816 | + const std::string policy_args = GetSimtPolicyArgs(plan, input_info, index_info, tpipe); | ||
| 817 | + if (!policy_args.empty()) { | ||
| 818 | + ss << ", " << policy_args; | ||
| 819 | + } | ||
| 820 | + ss << ");" << std::endl; | ||
| 821 | +} | ||
| 822 | + | ||
| 631 | af::Status GenerateSimtContextInitializer(const std::string &context_name, const std::vector<SimtGmTensor> &gm_tensors, | 823 | af::Status GenerateSimtContextInitializer(const std::string &context_name, const std::vector<SimtGmTensor> &gm_tensors, |
| 632 | - std::stringstream &ss) { | 824 | + const TPipe &tpipe, std::stringstream &ss) { |
| 633 | ss << " " << context_name << " context{"; | 825 | ss << " " << context_name << " context{"; |
| 634 | for (size_t i = 0UL; i < gm_tensors.size(); ++i) { | 826 | for (size_t i = 0UL; i < gm_tensors.size(); ++i) { |
| 635 | const SimtGmTensor &gm_tensor = gm_tensors[i]; | 827 | const SimtGmTensor &gm_tensor = gm_tensors[i]; |
| 636 | std::string dtype; | 828 | std::string dtype; |
| 637 | GE_ASSERT_SUCCESS(Tensor::DtypeName(gm_tensor.dtype, dtype)); | 829 | GE_ASSERT_SUCCESS(Tensor::DtypeName(gm_tensor.dtype, dtype)); |
| 638 | - ss << (i == 0UL ? "" : ", ") << "(__gm__ " << dtype << " *)" << kGlobalTensorNamePrefix << gm_tensor.gm_tensor_id | 830 | + if (gm_tensor.is_scalar) { |
| 639 | - << ".GetPhyAddr()"; | 831 | + const Tensor *scalar = tpipe.GetTensor(gm_tensor.value_tensor_id); |
| 832 | + GE_ASSERT_NOTNULL(scalar, "IndirectLoad SIMT ScalarData tensor is missing."); | ||
| 833 | + ss << (i == 0UL ? "" : ", ") << scalar->name; | ||
| 834 | + } else { | ||
| 835 | + ss << (i == 0UL ? "" : ", ") << "(__gm__ " << dtype << " *)" << kGlobalTensorNamePrefix << gm_tensor.gm_tensor_id | ||
| 836 | + << ".GetPhyAddr()"; | ||
| 837 | + } | ||
| 640 | } | 838 | } |
| 641 | ss << "};" << std::endl; | 839 | ss << "};" << std::endl; |
| 642 | return af::SUCCESS; | 840 | return af::SUCCESS; |
| @@ -746,13 +944,7 @@ Status IndirectLoadRegApiCall::GenerateFuncDefinition(const TPipe &tpipe, const | |||
| 746 | node_name.c_str(), input_info.sizes.size(), axis_, index_nodes_.size(), output_nodes_.size(), | 944 | node_name.c_str(), input_info.sizes.size(), axis_, index_nodes_.size(), output_nodes_.size(), |
| 747 | simt_gm_tensors_.size()); | 945 | simt_gm_tensors_.size()); |
| 748 | 946 | ||
| 749 | - ss << "struct " << context_name << " {" << std::endl; | 947 | + GE_ASSERT_SUCCESS(GenerateSimtContextDefinition(context_name, simt_gm_tensors_, ss)); |
| 750 | - for (const SimtGmTensor &tensor : simt_gm_tensors_) { | ||
| 751 | - std::string dtype; | ||
| 752 | - GE_ASSERT_SUCCESS(Tensor::DtypeName(tensor.dtype, dtype)); | ||
| 753 | - ss << " __gm__ " << dtype << " *" << kSimtGmFieldNamePrefix << tensor.value_tensor_id << ";" << std::endl; | ||
| 754 | - } | ||
| 755 | - ss << "};" << std::endl; | ||
| 756 | ss << "struct " << body_name << " {" << std::endl; | 948 | ss << "struct " << body_name << " {" << std::endl; |
| 757 | ss << " using Context = " << context_name << ";" << std::endl; | 949 | ss << " using Context = " << context_name << ";" << std::endl; |
| 758 | SimtNodeSet index_node_set; | 950 | SimtNodeSet index_node_set; |
| @@ -848,11 +1040,9 @@ Status IndirectLoadRegApiCall::GenerateSimd(const TPipe &tpipe, const std::vecto | |||
| 848 | LogicalTensorInfo index_info; | 1040 | LogicalTensorInfo index_info; |
| 849 | GE_ASSERT_SUCCESS(BuildTensorWindowInfo(logical_view_.input, input, axis_pos, input_info)); | 1041 | GE_ASSERT_SUCCESS(BuildTensorWindowInfo(logical_view_.input, input, axis_pos, input_info)); |
| 850 | GE_ASSERT_SUCCESS(BuildTensorWindowInfo(logical_view_.index, index, axis_pos, index_info)); | 1042 | GE_ASSERT_SUCCESS(BuildTensorWindowInfo(logical_view_.index, index, axis_pos, index_info)); |
| 851 | - const bool requires_strided_api = | 1043 | + const SimdApiKind api_kind = SelectSimdApi(logical_view_, implementation_); |
| 852 | - logical_view_.input.kind != ascgen_utils::indirect_load::IndirectLoadLayoutKind::kDense || | ||
| 853 | - logical_view_.index.kind != ascgen_utils::indirect_load::IndirectLoadLayoutKind::kDense; | ||
| 854 | const auto tmp_iter = tmp_buf_id.find(-1L); | 1044 | const auto tmp_iter = tmp_buf_id.find(-1L); |
| 855 | - if (requires_strided_api) { | 1045 | + if (api_kind == SimdApiKind::kStrided) { |
| 856 | GE_ASSERT_TRUE(tmp_iter != tmp_buf_id.end(), "IndirectLoad SIMD requires an API-level tmp buffer."); | 1046 | GE_ASSERT_TRUE(tmp_iter != tmp_buf_id.end(), "IndirectLoad SIMD requires an API-level tmp buffer."); |
| 857 | } | 1047 | } |
| 858 | 1048 | ||
| @@ -868,23 +1058,14 @@ Status IndirectLoadRegApiCall::GenerateSimd(const TPipe &tpipe, const std::vecto | |||
| 868 | std::stringstream ss; | 1058 | std::stringstream ss; |
| 869 | ss << "// IndirectLoad SIMD" << std::endl; | 1059 | ss << "// IndirectLoad SIMD" << std::endl; |
| 870 | ss << "{" << std::endl; | 1060 | ss << "{" << std::endl; |
| 871 | - const char *api = !requires_strided_api && implementation_ == ascgen_utils::indirect_load::Implementation::kGatherApi | 1061 | + if (api_kind == SimdApiKind::kStrided) { |
| 872 | - ? "IndirectLoadSimdGatherApi" | 1062 | + EmitSimdStridedParams(tpipe, current_axis, output, input_info, index_info, logical_view_.output, |
| 873 | - : "IndirectLoadSimd"; | 1063 | + input_info.sizes.size(), ss); |
| 874 | - ss << " AscendC::" << api << "<" << input_dtype << ", " << index_dtype << ", " << input_info.sizes.size() << ", " | ||
| 875 | - << axis_ << ">(" << std::endl; | ||
| 876 | - ss << " " << input << ", " << index << ", " << output << ", "; | ||
| 877 | - if (requires_strided_api) { | ||
| 878 | - ss << tpipe.tmp_buf.name << "_" << tmp_iter->second << ", " << output.actual_size << ", " | ||
| 879 | - << tpipe.tiler.Offset(current_axis, output.axis, output.axis_strides) << ", " | ||
| 880 | - << JoinSizeExprs(index_info.sizes, tpipe) << ", " << JoinSizeExprs(input_info.strides, tpipe) << ", " | ||
| 881 | - << JoinSizeExprs(index_info.strides, tpipe); | ||
| 882 | - } else { | ||
| 883 | - ss << output.actual_size << ", " << tpipe.tiler.Offset(current_axis, output.axis, output.axis_strides) << ", " | ||
| 884 | - << input.actual_size << ", " << tpipe.tiler.Size(input_info.sizes[axis_pos]) << ", " | ||
| 885 | - << JoinSizeExprs(index_info.sizes, tpipe) << ", " << JoinSizeExprs(input_info.strides, tpipe); | ||
| 886 | } | 1064 | } |
| 887 | - ss << ");" << std::endl; | 1065 | + const std::string tmp_name = |
| 1066 | + api_kind == SimdApiKind::kStrided ? tpipe.tmp_buf.name + "_" + std::to_string(tmp_iter->second) : ""; | ||
| 1067 | + EmitSimdInvocation(tpipe, current_axis, input, index, output, input_info, index_info, axis_, axis_pos, api_kind, | ||
| 1068 | + input_dtype, index_dtype, tmp_name, ss); | ||
| 888 | ss << "}" << std::endl; | 1069 | ss << "}" << std::endl; |
| 889 | result = ss.str(); | 1070 | result = ss.str(); |
| 890 | return af::SUCCESS; | 1071 | return af::SUCCESS; |
| @@ -909,23 +1090,8 @@ Status IndirectLoadRegApiCall::GenerateSimtInvocation(const TPipe &tpipe, const | |||
| 909 | GE_ASSERT_NOTNULL(output_tensor, "IndirectLoad SIMT UB output tensor is missing."); | 1090 | GE_ASSERT_NOTNULL(output_tensor, "IndirectLoad SIMT UB output tensor is missing."); |
| 910 | GE_ASSERT_SUCCESS(CalcVectorizedElementCount(*output_tensor, output_element_count)); | 1091 | GE_ASSERT_SUCCESS(CalcVectorizedElementCount(*output_tensor, output_element_count)); |
| 911 | } | 1092 | } |
| 912 | - ss << " AscendC::IndirectLoadSimt<" << input_dtype << ", " << output_dtype << ", "; | 1093 | + EmitSimtInvocation(tpipe, input_dtype, output_dtype, outer_tb_var, body_name, input_info, index_info, plan, axis_, |
| 913 | - ss << body_name << ", " << GetSimtPolicyType(plan, input_info.sizes.size(), axis_); | 1094 | + has_post_reduce_, output_tensor, output_element_count, ss); |
| 914 | - ss << ">(" << std::endl; | ||
| 915 | - if (has_post_reduce_) { | ||
| 916 | - const std::string output_elements = tpipe.tiler.Size(output_element_count); | ||
| 917 | - ss << " input_ptr, " << *output_tensor << ", context, static_cast<uint32_t>(" << output_elements << "), "; | ||
| 918 | - ss << "(static_cast<" << plan.offset_type << ">(block_dim_offset) + static_cast<" << plan.offset_type << ">("; | ||
| 919 | - ss << outer_tb_var << ")) * " << PromoteSizeExpr(output_elements, plan.offset_type); | ||
| 920 | - } else { | ||
| 921 | - ss << " input_ptr, y_ptr, context, static_cast<uint32_t>(" << outer_tb_var << "_loop_size), "; | ||
| 922 | - ss << "static_cast<" << plan.offset_type << ">(block_dim_offset)"; | ||
| 923 | - } | ||
| 924 | - const std::string policy_args = GetSimtPolicyArgs(plan, input_info, index_info, tpipe); | ||
| 925 | - if (!policy_args.empty()) { | ||
| 926 | - ss << ", " << policy_args; | ||
| 927 | - } | ||
| 928 | - ss << ");" << std::endl; | ||
| 929 | return af::SUCCESS; | 1095 | return af::SUCCESS; |
| 930 | } | 1096 | } |
| 931 | 1097 | ||
| @@ -956,7 +1122,7 @@ Status IndirectLoadRegApiCall::GenerateSimt(const TPipe &tpipe, const std::vecto | |||
| 956 | ss << " __gm__ " << output_dtype << " *y_ptr = (__gm__ " << output_dtype << " *)" << output_gm_tensor_ | 1122 | ss << " __gm__ " << output_dtype << " *y_ptr = (__gm__ " << output_dtype << " *)" << output_gm_tensor_ |
| 957 | << ".GetPhyAddr();" << std::endl; | 1123 | << ".GetPhyAddr();" << std::endl; |
| 958 | } | 1124 | } |
| 959 | - GE_ASSERT_SUCCESS(GenerateSimtContextInitializer(context_name, simt_gm_tensors_, ss)); | 1125 | + GE_ASSERT_SUCCESS(GenerateSimtContextInitializer(context_name, simt_gm_tensors_, tpipe, ss)); |
| 960 | GE_ASSERT_SUCCESS(GenerateSimtInvocation(tpipe, input_dtype, output_dtype, outer_tb_var, ss)); | 1126 | GE_ASSERT_SUCCESS(GenerateSimtInvocation(tpipe, input_dtype, output_dtype, outer_tb_var, ss)); |
| 961 | ss << "}" << std::endl; | 1127 | ss << "}" << std::endl; |
| 962 | result = ss.str(); | 1128 | result = ss.str(); |
| @@ -22,6 +22,7 @@ struct SimtGmTensor { | |||
| 22 | ascir::TensorId value_tensor_id; | 22 | ascir::TensorId value_tensor_id; |
| 23 | ascir::TensorId gm_tensor_id; | 23 | ascir::TensorId gm_tensor_id; |
| 24 | af::DataType dtype; | 24 | af::DataType dtype; |
| 25 | + bool is_scalar{false}; | ||
| 25 | }; | 26 | }; |
| 26 | 27 | ||
| 27 | class IndirectLoadRegApiCall final : public ApiCall { | 28 | class IndirectLoadRegApiCall final : public ApiCall { |
| @@ -692,6 +692,7 @@ build_backend() { | |||
| 692 | cmake $CMAKE_ARGS ../ | 692 | cmake $CMAKE_ARGS ../ |
| 693 | 693 | ||
| 694 | # st用例可执行文件的列表,inductor split_compile 仅保留 presubmit 代表用例,其余放 nightly。 | 694 | # st用例可执行文件的列表,inductor split_compile 仅保留 presubmit 代表用例,其余放 nightly。 |
| 695 | + # Dual-IL case is kept in CMake for local debugging but excluded from the online pipeline. | ||
| 695 | MAKE_TARGET_LIST="add_abs_test_e2e \ | 696 | MAKE_TARGET_LIST="add_abs_test_e2e \ |
| 696 | axpy_abs_test_e2e \ | 697 | axpy_abs_test_e2e \ |
| 697 | sub_abs_test_e2e \ | 698 | sub_abs_test_e2e \ |
| @@ -826,6 +827,24 @@ build_backend() { | |||
| 826 | indirect_load_rank3_axis1_input_index_gap_sk_test_e2e_v2 \ | 827 | indirect_load_rank3_axis1_input_index_gap_sk_test_e2e_v2 \ |
| 827 | indirect_load_rank3_axis1_input_index_outer_gap_simt_test_e2e_v2 \ | 828 | indirect_load_rank3_axis1_input_index_outer_gap_simt_test_e2e_v2 \ |
| 828 | indirect_load_rank3_axis1_torch_gather_frontend_e2e_v2 \ | 829 | indirect_load_rank3_axis1_torch_gather_frontend_e2e_v2 \ |
| 830 | + indirect_load_graph_hint_reduce_simt_test_e2e_v2 \ | ||
| 831 | + indirect_load_graph_hint_simd_repro_e2e_v2 \ | ||
| 832 | + indirect_load_embedding_reduce_simt_test_e2e_v2 \ | ||
| 833 | + indirect_load_add_il_reduce_test_e2e_v2 \ | ||
| 834 | + indirect_load_user_embedding_sum_e2e_v2 \ | ||
| 835 | + indirect_load_user_embedding_mul_e2e_v2 \ | ||
| 836 | + indirect_load_user_layernorm_e2e_v2 \ | ||
| 837 | + indirect_load_user_layernorm_simd_e2e_v2 \ | ||
| 838 | + indirect_load_user_embedding_exp_abs_add_simd_e2e_v2 \ | ||
| 839 | + indirect_load_user_embedding_exp_abs_add_simt_e2e_v2 \ | ||
| 840 | + indirect_load_user_fanout_direct_stores_simd_e2e_v2 \ | ||
| 841 | + indirect_load_user_fanout_direct_stores_simt_e2e_v2 \ | ||
| 842 | + indirect_load_user_fanout_direct_reduce_simd_e2e_v2 \ | ||
| 843 | + indirect_load_user_fanout_direct_reduce_simt_e2e_v2 \ | ||
| 844 | + indirect_load_user_fanout_post_stores_simd_e2e_v2 \ | ||
| 845 | + indirect_load_user_fanout_post_stores_simt_e2e_v2 \ | ||
| 846 | + indirect_load_user_fanout_post_reduce_simd_e2e_v2 \ | ||
| 847 | + indirect_load_user_fanout_post_reduce_simt_e2e_v2 \ | ||
| 829 | load_where_x2_x3_is_ubscalar_store_test_e2e_v2 \ | 848 | load_where_x2_x3_is_ubscalar_store_test_e2e_v2 \ |
| 830 | gather_reduce_store_test_e2e_v2 \ | 849 | gather_reduce_store_test_e2e_v2 \ |
| 831 | load_where_store_test_e2e_v2 \ | 850 | load_where_store_test_e2e_v2 \ |