forked from huawei/mindspore2022
!16894 clean codex
From: @zhupuxu Reviewed-by: @jjfeing,@kisnwang Signed-off-by: @kisnwang
This commit is contained in:
commit
22d62f20ff
|
|
@ -279,6 +279,7 @@ const AnfNodePtr AdamApplyOneWithDecayRule::Process(const FuncGraphPtr &graph, c
|
|||
return nullptr;
|
||||
}
|
||||
auto sub0 = node;
|
||||
constexpr size_t kOutputIndex2 = 2;
|
||||
if (AnfAlgo::CheckPrimitiveType(node, prim::kPrimDepend)) {
|
||||
auto iter_sub0 = (*equiv).find(sub0_var_);
|
||||
if (iter_sub0 == (*equiv).end()) {
|
||||
|
|
@ -327,7 +328,7 @@ const AnfNodePtr AdamApplyOneWithDecayRule::Process(const FuncGraphPtr &graph, c
|
|||
MS_EXCEPTION_IF_NULL(manager);
|
||||
(void)manager->Replace(add1, fusion_node_outputs[0]);
|
||||
(void)manager->Replace(add0, fusion_node_outputs[1]);
|
||||
return fusion_node_outputs[2];
|
||||
return fusion_node_outputs[kOutputIndex2];
|
||||
}
|
||||
} // namespace opt
|
||||
} // namespace mindspore
|
||||
|
|
|
|||
|
|
@ -31,9 +31,12 @@ CNodePtr CreateBNInferGrad(const FuncGraphPtr &graph, const CNodePtr &batchnormg
|
|||
MS_EXCEPTION_IF_NULL(batchnormgrad);
|
||||
auto prim = std::make_shared<Primitive>(kBNInferGradOpName);
|
||||
std::vector<AnfNodePtr> inputs = {NewValueNode(prim)};
|
||||
inputs.push_back(batchnormgrad->input(1));
|
||||
inputs.push_back(batchnormgrad->input(3));
|
||||
inputs.push_back(batchnormgrad->input(5));
|
||||
constexpr size_t kDBatchMean = 1;
|
||||
constexpr size_t kInputX = 3;
|
||||
constexpr size_t kBatchStd = 5;
|
||||
inputs.push_back(batchnormgrad->input(kDBatchMean));
|
||||
inputs.push_back(batchnormgrad->input(kInputX));
|
||||
inputs.push_back(batchnormgrad->input(kBatchStd));
|
||||
auto new_node = graph->NewCNode(inputs);
|
||||
MS_EXCEPTION_IF_NULL(new_node);
|
||||
new_node->set_scope(batchnormgrad->scope());
|
||||
|
|
|
|||
|
|
@ -30,8 +30,9 @@ const size_t kReluV2OutputNum = 2;
|
|||
|
||||
CNodePtr GetRelu(const CNodePtr &relu_grad) {
|
||||
MS_EXCEPTION_IF_NULL(relu_grad);
|
||||
constexpr size_t kReluIndex2 = 2;
|
||||
CheckCNodeInputSize(relu_grad, kReluGradInputTensorNum);
|
||||
auto relu_anf = relu_grad->input(2);
|
||||
auto relu_anf = relu_grad->input(kReluIndex2);
|
||||
MS_EXCEPTION_IF_NULL(relu_anf);
|
||||
return relu_anf->cast<CNodePtr>();
|
||||
}
|
||||
|
|
@ -40,7 +41,7 @@ CNodePtr CreateReluV2(const FuncGraphPtr &graph, const CNodePtr &relu) {
|
|||
MS_EXCEPTION_IF_NULL(graph);
|
||||
MS_EXCEPTION_IF_NULL(relu);
|
||||
CheckCNodeInputSize(relu, kReluInputTensorNum);
|
||||
|
||||
constexpr auto kMaskShapeSize = 4;
|
||||
auto prim = std::make_shared<Primitive>(kReluV2OpName);
|
||||
std::vector<AnfNodePtr> inputs = {NewValueNode(prim), relu->input(1)};
|
||||
auto new_node = graph->NewCNode(inputs);
|
||||
|
|
@ -53,17 +54,24 @@ CNodePtr CreateReluV2(const FuncGraphPtr &graph, const CNodePtr &relu) {
|
|||
return nullptr;
|
||||
}
|
||||
std::vector<size_t> mask_shape = AnfAlgo::GetOutputInferShape(relu, 0);
|
||||
if (mask_shape.size() != 4) {
|
||||
if (mask_shape.size() != kMaskShapeSize) {
|
||||
MS_LOG(DEBUG) << "relu's infer shape size not equal 4";
|
||||
return nullptr;
|
||||
}
|
||||
auto input_dtype = AnfAlgo::GetPrevNodeOutputInferDataType(relu, 0);
|
||||
constexpr auto kMultiplierInt8 = 31;
|
||||
constexpr auto kDivisorInt8 = 32;
|
||||
constexpr auto kMultiplier = 15;
|
||||
constexpr auto kDivisor = 16;
|
||||
constexpr auto kThirdShapeInt8 = 4;
|
||||
constexpr auto kThirdShape = 2;
|
||||
|
||||
if (input_dtype == kNumberTypeUInt8 || input_dtype == kNumberTypeInt8) {
|
||||
mask_shape[1] = (mask_shape[1] + 31) / 32;
|
||||
mask_shape.push_back(4);
|
||||
mask_shape[1] = (mask_shape[1] + kMultiplierInt8) / kDivisorInt8;
|
||||
mask_shape.push_back(kThirdShapeInt8);
|
||||
} else {
|
||||
mask_shape[1] = (mask_shape[1] + 15) / 16;
|
||||
mask_shape.push_back(2);
|
||||
mask_shape[1] = (mask_shape[1] + kMultiplier) / kDivisor;
|
||||
mask_shape.push_back(kThirdShape);
|
||||
}
|
||||
|
||||
auto types = {AnfAlgo::GetOutputInferDataType(relu, 0), mask_dtype};
|
||||
|
|
|
|||
|
|
@ -138,8 +138,11 @@ void FusedBatchNormFusion::GetBNTrainingUpdateAbstractList(const EquivPtr &equiv
|
|||
auto variable_input1 = GetAnfNodeByVar(equiv, variable_input1_var_);
|
||||
MS_EXCEPTION_IF_NULL(variable_input0);
|
||||
MS_EXCEPTION_IF_NULL(variable_input1);
|
||||
*abstract_list = {bn_abstract_tuple->elements()[0], variable_input0->abstract(), variable_input1->abstract(),
|
||||
bn_abstract_tuple->elements()[1], bn_abstract_tuple->elements()[2]};
|
||||
constexpr size_t kElements0 = 0;
|
||||
constexpr size_t kElements1 = 1;
|
||||
constexpr size_t kElements2 = 2;
|
||||
*abstract_list = {bn_abstract_tuple->elements()[kElements0], variable_input0->abstract(), variable_input1->abstract(),
|
||||
bn_abstract_tuple->elements()[kElements1], bn_abstract_tuple->elements()[kElements2]};
|
||||
}
|
||||
|
||||
AnfNodePtr FusedBatchNormFusion::CreateBNTrainingUpdate(
|
||||
|
|
|
|||
|
|
@ -30,11 +30,12 @@ bool LambNextMVRule::IsRuleMatched(const FuncGraphPtr &func_graph, const AnfNode
|
|||
MS_EXCEPTION_IF_NULL(equiv);
|
||||
auto real_div0 = GetAnfNodeByVar(equiv, real_div0_var_);
|
||||
auto real_div2 = GetAnfNodeByVar(equiv, real_div2_var_);
|
||||
constexpr size_t kRealDiv0Size = 2;
|
||||
|
||||
auto manager = func_graph->manager();
|
||||
MS_EXCEPTION_IF_NULL(manager);
|
||||
auto &users = manager->node_users();
|
||||
if (users.find(real_div0) == users.end() || users[real_div0].size() < 2) {
|
||||
if (users.find(real_div0) == users.end() || users[real_div0].size() < kRealDiv0Size) {
|
||||
return false;
|
||||
}
|
||||
AnfNodeIndexSet real_div0_outputs = users[real_div0];
|
||||
|
|
@ -60,6 +61,10 @@ AnfNodePtr LambNextMVRule::CreateLambNextMVNode(const FuncGraphPtr &func_graph,
|
|||
const EquivPtr &equiv) const {
|
||||
MS_EXCEPTION_IF_NULL(func_graph);
|
||||
auto prim = std::make_shared<Primitive>(kLambNextMVOpName);
|
||||
constexpr size_t kOutputsIndex1 = 1;
|
||||
constexpr size_t kOutputsIndex2 = 2;
|
||||
constexpr size_t kOutputsIndex3 = 3;
|
||||
|
||||
std::vector<AnfNodePtr> lamb_next_mv_rule_inputs = {NewValueNode(prim)};
|
||||
lamb_next_mv_rule_inputs.push_back(utils::cast<AnfNodePtr>((*equiv)[input0_]));
|
||||
lamb_next_mv_rule_inputs.push_back(utils::cast<AnfNodePtr>((*equiv)[input1_]));
|
||||
|
|
@ -91,9 +96,9 @@ AnfNodePtr LambNextMVRule::CreateLambNextMVNode(const FuncGraphPtr &func_graph,
|
|||
|
||||
auto manager = func_graph->manager();
|
||||
MS_EXCEPTION_IF_NULL(manager);
|
||||
(void)manager->Replace(old_pattern_outputs[1], lamb_next_mv_rule_outputs[1]);
|
||||
(void)manager->Replace(old_pattern_outputs[2], lamb_next_mv_rule_outputs[2]);
|
||||
(void)manager->Replace(old_pattern_outputs[3], lamb_next_mv_rule_outputs[3]);
|
||||
(void)manager->Replace(old_pattern_outputs[kOutputsIndex1], lamb_next_mv_rule_outputs[kOutputsIndex1]);
|
||||
(void)manager->Replace(old_pattern_outputs[kOutputsIndex2], lamb_next_mv_rule_outputs[kOutputsIndex2]);
|
||||
(void)manager->Replace(old_pattern_outputs[kOutputsIndex3], lamb_next_mv_rule_outputs[kOutputsIndex3]);
|
||||
|
||||
return lamb_next_mv_rule_outputs[0];
|
||||
}
|
||||
|
|
|
|||
|
|
@ -19,6 +19,16 @@
|
|||
#include "frontend/optimizer/opt.h"
|
||||
#include "utils/trace_base.h"
|
||||
namespace mindspore {
|
||||
namespace {
|
||||
constexpr size_t kZeroIndex = 0;
|
||||
constexpr size_t kFirstIndex = 1;
|
||||
constexpr size_t kSecondIndex = 2;
|
||||
constexpr size_t kThirdIndex = 3;
|
||||
constexpr size_t kFourthIndex = 4;
|
||||
constexpr size_t kFifthIndex = 5;
|
||||
constexpr size_t kSixthIndex = 6;
|
||||
} // namespace
|
||||
|
||||
namespace opt {
|
||||
AnfNodePtr LambNextMVWithDecayRule::GetLambNextMVWithDecayOutput(const FuncGraphPtr &func_graph,
|
||||
const AnfNodePtr &new_node, const AnfNodePtr &add3,
|
||||
|
|
@ -136,18 +146,18 @@ const BaseRef LambNextMVWithDecayRuleCond1::DefinePattern() const {
|
|||
MS_EXCEPTION_IF_NULL(prim_sqrt);
|
||||
const auto prim_deal_div = std::make_shared<Primitive>(kRealDivOpName);
|
||||
MS_EXCEPTION_IF_NULL(prim_deal_div);
|
||||
VectorRef mul2 = VectorRef({prim::kPrimMul, input_vars_[1], constant_mul_input_vars_[2]});
|
||||
VectorRef mul3 = VectorRef({prim::kPrimMul, input_vars_[0], constant_mul_input_vars_[3]});
|
||||
VectorRef mul2 = VectorRef({prim::kPrimMul, input_vars_[kFirstIndex], constant_mul_input_vars_[kSecondIndex]});
|
||||
VectorRef mul3 = VectorRef({prim::kPrimMul, input_vars_[kZeroIndex], constant_mul_input_vars_[kThirdIndex]});
|
||||
VectorRef add1 = VectorRef({add1_var_, mul2, mul3});
|
||||
VectorRef real_div1 = VectorRef({real_div1_var_, add1, input_vars_[2]});
|
||||
VectorRef real_div1 = VectorRef({real_div1_var_, add1, input_vars_[kSecondIndex]});
|
||||
VectorRef sqrt1 = VectorRef({prim_sqrt, real_div1});
|
||||
VectorRef add4 = VectorRef({prim::kPrimAdd, sqrt1, constant_add2_y_});
|
||||
VectorRef mul0 = VectorRef({prim::kPrimMul, input_vars_[4], constant_mul_input_vars_[0]});
|
||||
VectorRef mul1 = VectorRef({prim::kPrimMul, input_vars_[3], constant_mul_input_vars_[1]});
|
||||
VectorRef mul0 = VectorRef({prim::kPrimMul, input_vars_[kFourthIndex], constant_mul_input_vars_[kZeroIndex]});
|
||||
VectorRef mul1 = VectorRef({prim::kPrimMul, input_vars_[kThirdIndex], constant_mul_input_vars_[kFirstIndex]});
|
||||
VectorRef add0 = VectorRef({add0_var_, mul0, mul1});
|
||||
VectorRef real_div0 = VectorRef({real_div0_var_, add0, input_vars_[5]});
|
||||
VectorRef real_div0 = VectorRef({real_div0_var_, add0, input_vars_[kFifthIndex]});
|
||||
VectorRef real_div4 = VectorRef({prim_deal_div, real_div0, add4});
|
||||
VectorRef mul4 = VectorRef({mul4_var_, constant_mul_input_vars_[4], input_vars_[6]});
|
||||
VectorRef mul4 = VectorRef({mul4_var_, constant_mul_input_vars_[kFourthIndex], input_vars_[kSixthIndex]});
|
||||
VectorRef add5 = VectorRef({prim::kPrimAdd, mul4, real_div4});
|
||||
return add5;
|
||||
}
|
||||
|
|
@ -177,18 +187,18 @@ const BaseRef LambNextMVWithDecayRuleCond2::DefinePattern() const {
|
|||
MS_EXCEPTION_IF_NULL(prim_sqrt);
|
||||
const auto prim_deal_div = std::make_shared<Primitive>(kRealDivOpName);
|
||||
MS_EXCEPTION_IF_NULL(prim_deal_div);
|
||||
VectorRef mul2 = VectorRef({prim::kPrimMul, constant_mul_input_vars_[2], input_vars_[1]});
|
||||
VectorRef mul3 = VectorRef({prim::kPrimMul, constant_mul_input_vars_[3], input_vars_[0]});
|
||||
VectorRef mul2 = VectorRef({prim::kPrimMul, constant_mul_input_vars_[kSecondIndex], input_vars_[kFirstIndex]});
|
||||
VectorRef mul3 = VectorRef({prim::kPrimMul, constant_mul_input_vars_[kThirdIndex], input_vars_[kZeroIndex]});
|
||||
VectorRef add1 = VectorRef({add1_var_, mul2, mul3});
|
||||
VectorRef real_div1 = VectorRef({real_div1_var_, add1, input_vars_[2]});
|
||||
VectorRef real_div1 = VectorRef({real_div1_var_, add1, input_vars_[kSecondIndex]});
|
||||
VectorRef sqrt1 = VectorRef({prim_sqrt, real_div1});
|
||||
VectorRef add4 = VectorRef({prim::kPrimAdd, constant_add2_y_, sqrt1});
|
||||
VectorRef mul0 = VectorRef({prim::kPrimMul, constant_mul_input_vars_[0], input_vars_[4]});
|
||||
VectorRef mul1 = VectorRef({prim::kPrimMul, constant_mul_input_vars_[1], input_vars_[3]});
|
||||
VectorRef mul0 = VectorRef({prim::kPrimMul, constant_mul_input_vars_[kZeroIndex], input_vars_[kFourthIndex]});
|
||||
VectorRef mul1 = VectorRef({prim::kPrimMul, constant_mul_input_vars_[kFirstIndex], input_vars_[kThirdIndex]});
|
||||
VectorRef add0 = VectorRef({add0_var_, mul0, mul1});
|
||||
VectorRef real_div0 = VectorRef({real_div0_var_, add0, input_vars_[5]});
|
||||
VectorRef real_div0 = VectorRef({real_div0_var_, add0, input_vars_[kFifthIndex]});
|
||||
VectorRef real_div4 = VectorRef({prim_deal_div, real_div0, add4});
|
||||
VectorRef mul4 = VectorRef({mul4_var_, constant_mul_input_vars_[4], input_vars_[6]});
|
||||
VectorRef mul4 = VectorRef({mul4_var_, constant_mul_input_vars_[kFourthIndex], input_vars_[kSixthIndex]});
|
||||
VectorRef add5 = VectorRef({prim::kPrimAdd, mul4, real_div4});
|
||||
return add5;
|
||||
}
|
||||
|
|
@ -218,18 +228,18 @@ const BaseRef LambNextMVWithDecayRuleCond3::DefinePattern() const {
|
|||
MS_EXCEPTION_IF_NULL(prim_sqrt);
|
||||
const auto prim_deal_div = std::make_shared<Primitive>(kRealDivOpName);
|
||||
MS_EXCEPTION_IF_NULL(prim_deal_div);
|
||||
VectorRef mul2 = VectorRef({prim::kPrimMul, input_vars_[1], constant_mul_input_vars_[2]});
|
||||
VectorRef mul3 = VectorRef({prim::kPrimMul, constant_mul_input_vars_[3], input_vars_[0]});
|
||||
VectorRef mul2 = VectorRef({prim::kPrimMul, input_vars_[kFirstIndex], constant_mul_input_vars_[kSecondIndex]});
|
||||
VectorRef mul3 = VectorRef({prim::kPrimMul, constant_mul_input_vars_[kThirdIndex], input_vars_[kZeroIndex]});
|
||||
VectorRef add1 = VectorRef({add1_var_, mul2, mul3});
|
||||
VectorRef real_div1 = VectorRef({real_div1_var_, add1, input_vars_[2]});
|
||||
VectorRef real_div1 = VectorRef({real_div1_var_, add1, input_vars_[kSecondIndex]});
|
||||
VectorRef sqrt1 = VectorRef({prim_sqrt, real_div1});
|
||||
VectorRef add4 = VectorRef({prim::kPrimAdd, sqrt1, constant_add2_y_});
|
||||
VectorRef mul0 = VectorRef({prim::kPrimMul, input_vars_[4], constant_mul_input_vars_[0]});
|
||||
VectorRef mul1 = VectorRef({prim::kPrimMul, input_vars_[3], constant_mul_input_vars_[1]});
|
||||
VectorRef mul0 = VectorRef({prim::kPrimMul, input_vars_[kFourthIndex], constant_mul_input_vars_[kZeroIndex]});
|
||||
VectorRef mul1 = VectorRef({prim::kPrimMul, input_vars_[kThirdIndex], constant_mul_input_vars_[kFirstIndex]});
|
||||
VectorRef add0 = VectorRef({add0_var_, mul0, mul1});
|
||||
VectorRef real_div0 = VectorRef({real_div0_var_, add0, input_vars_[5]});
|
||||
VectorRef real_div0 = VectorRef({real_div0_var_, add0, input_vars_[kFifthIndex]});
|
||||
VectorRef real_div4 = VectorRef({prim_deal_div, real_div0, add4});
|
||||
VectorRef mul4 = VectorRef({mul4_var_, input_vars_[6], constant_mul_input_vars_[4]});
|
||||
VectorRef mul4 = VectorRef({mul4_var_, input_vars_[kSixthIndex], constant_mul_input_vars_[kFourthIndex]});
|
||||
VectorRef add5 = VectorRef({prim::kPrimAdd, mul4, real_div4});
|
||||
return add5;
|
||||
}
|
||||
|
|
@ -260,18 +270,18 @@ const BaseRef LambNextMVWithDecayRuleCond4::DefinePattern() const {
|
|||
MS_EXCEPTION_IF_NULL(prim_sqrt);
|
||||
const auto prim_deal_div = std::make_shared<Primitive>(kRealDivOpName);
|
||||
MS_EXCEPTION_IF_NULL(prim_deal_div);
|
||||
VectorRef mul2 = VectorRef({prim::kPrimMul, constant_mul_input_vars_[2], input_vars_[1]});
|
||||
VectorRef mul3 = VectorRef({prim::kPrimMul, constant_mul_input_vars_[3], input_vars_[0]});
|
||||
VectorRef mul2 = VectorRef({prim::kPrimMul, constant_mul_input_vars_[kSecondIndex], input_vars_[kFirstIndex]});
|
||||
VectorRef mul3 = VectorRef({prim::kPrimMul, constant_mul_input_vars_[kThirdIndex], input_vars_[kZeroIndex]});
|
||||
VectorRef add1 = VectorRef({add1_var_, mul2, mul3});
|
||||
VectorRef real_div1 = VectorRef({real_div1_var_, add1, input_vars_[2]});
|
||||
VectorRef real_div1 = VectorRef({real_div1_var_, add1, input_vars_[kSecondIndex]});
|
||||
VectorRef sqrt1 = VectorRef({prim_sqrt, real_div1});
|
||||
VectorRef add4 = VectorRef({prim::kPrimAdd, sqrt1, constant_add2_y_});
|
||||
VectorRef mul0 = VectorRef({prim::kPrimMul, constant_mul_input_vars_[0], input_vars_[4]});
|
||||
VectorRef mul1 = VectorRef({prim::kPrimMul, constant_mul_input_vars_[1], input_vars_[3]});
|
||||
VectorRef mul0 = VectorRef({prim::kPrimMul, constant_mul_input_vars_[kZeroIndex], input_vars_[kFourthIndex]});
|
||||
VectorRef mul1 = VectorRef({prim::kPrimMul, constant_mul_input_vars_[kFirstIndex], input_vars_[kThirdIndex]});
|
||||
VectorRef add0 = VectorRef({add0_var_, mul0, mul1});
|
||||
VectorRef real_div0 = VectorRef({real_div0_var_, add0, input_vars_[5]});
|
||||
VectorRef real_div0 = VectorRef({real_div0_var_, add0, input_vars_[kFifthIndex]});
|
||||
VectorRef real_div4 = VectorRef({prim_deal_div, real_div0, add4});
|
||||
VectorRef mul4 = VectorRef({mul4_var_, constant_mul_input_vars_[4], input_vars_[6]});
|
||||
VectorRef mul4 = VectorRef({mul4_var_, constant_mul_input_vars_[kFourthIndex], input_vars_[kSixthIndex]});
|
||||
VectorRef add5 = VectorRef({prim::kPrimAdd, real_div4, mul4});
|
||||
return add5;
|
||||
}
|
||||
|
|
|
|||
|
|
@ -26,26 +26,31 @@
|
|||
namespace mindspore {
|
||||
namespace opt {
|
||||
namespace {
|
||||
constexpr auto kFirstIndex1 = 1;
|
||||
constexpr auto kSecondIndex2 = 2;
|
||||
constexpr auto kThirdIndex3 = 3;
|
||||
|
||||
std::tuple<AnfNodePtr, AnfNodePtr, AnfNodePtr, AnfNodePtr> GetSharedNodes(const AnfNodePtr &node) {
|
||||
MS_EXCEPTION_IF_NULL(node);
|
||||
auto add3 = node->cast<CNodePtr>();
|
||||
MS_EXCEPTION_IF_NULL(add3);
|
||||
CheckCNodeInputSize(add3, kAddInputTensorNum);
|
||||
auto real_div2_anf = add3->input(1);
|
||||
auto real_div2_anf = add3->input(kFirstIndex1);
|
||||
MS_EXCEPTION_IF_NULL(real_div2_anf);
|
||||
auto real_div2 = real_div2_anf->cast<CNodePtr>();
|
||||
MS_EXCEPTION_IF_NULL(real_div2);
|
||||
CheckCNodeInputSize(real_div2, kRealDivInputTensorNum);
|
||||
auto sqrt0_anf = real_div2->input(2);
|
||||
auto sqrt0_anf = real_div2->input(kSecondIndex2);
|
||||
MS_EXCEPTION_IF_NULL(sqrt0_anf);
|
||||
auto sqrt0 = sqrt0_anf->cast<CNodePtr>();
|
||||
MS_EXCEPTION_IF_NULL(sqrt0);
|
||||
CheckCNodeInputSize(sqrt0, kSqrtInputTensorNum);
|
||||
auto add2_anf = sqrt0->input(1);
|
||||
auto add2_anf = sqrt0->input(kFirstIndex1);
|
||||
MS_EXCEPTION_IF_NULL(add2_anf);
|
||||
auto add2 = add2_anf->cast<CNodePtr>();
|
||||
CheckCNodeInputSize(add2, kAddInputTensorNum);
|
||||
return std::make_tuple(add3->input(2), real_div2->input(1), add2->input(1), add2->input(2));
|
||||
return std::make_tuple(add3->input(kSecondIndex2), real_div2->input(kFirstIndex1), add2->input(kFirstIndex1),
|
||||
add2->input(kSecondIndex2));
|
||||
}
|
||||
|
||||
bool MatchAdd5Pattern(const AnfNodePtr &node, const AnfNodePtr &mul4, const AnfNodePtr &real_div0,
|
||||
|
|
@ -57,7 +62,7 @@ bool MatchAdd5Pattern(const AnfNodePtr &node, const AnfNodePtr &mul4, const AnfN
|
|||
if (AnfAlgo::GetCNodeName(add5) != prim::kPrimAdd->name() || AnfAlgo::GetInputTensorNum(add5) != kAddInputTensorNum) {
|
||||
return false;
|
||||
}
|
||||
auto real_div4_anf = add5->input(1);
|
||||
auto real_div4_anf = add5->input(kFirstIndex1);
|
||||
if (real_div4_anf == nullptr || !real_div4_anf->isa<CNode>()) {
|
||||
return false;
|
||||
}
|
||||
|
|
@ -66,7 +71,7 @@ bool MatchAdd5Pattern(const AnfNodePtr &node, const AnfNodePtr &mul4, const AnfN
|
|||
AnfAlgo::GetInputTensorNum(real_div4) != kRealDivInputTensorNum) {
|
||||
return false;
|
||||
}
|
||||
auto add4_anf = real_div4->input(2);
|
||||
auto add4_anf = real_div4->input(kSecondIndex2);
|
||||
if (add4_anf == nullptr || !add4_anf->isa<CNode>()) {
|
||||
return false;
|
||||
}
|
||||
|
|
@ -74,7 +79,7 @@ bool MatchAdd5Pattern(const AnfNodePtr &node, const AnfNodePtr &mul4, const AnfN
|
|||
if (AnfAlgo::GetCNodeName(add4) != prim::kPrimAdd->name() || AnfAlgo::GetInputTensorNum(add4) != kAddInputTensorNum) {
|
||||
return false;
|
||||
}
|
||||
auto sqrt1_anf = add4->input(1);
|
||||
auto sqrt1_anf = add4->input(kFirstIndex1);
|
||||
if (sqrt1_anf == nullptr || !sqrt1_anf->isa<CNode>()) {
|
||||
return false;
|
||||
}
|
||||
|
|
@ -82,8 +87,8 @@ bool MatchAdd5Pattern(const AnfNodePtr &node, const AnfNodePtr &mul4, const AnfN
|
|||
if (AnfAlgo::GetCNodeName(sqrt1) != kSqrtOpName || AnfAlgo::GetInputTensorNum(sqrt1) != kSqrtInputTensorNum) {
|
||||
return false;
|
||||
}
|
||||
return add5->input(2) == mul4 && real_div4->input(1) == real_div0 && sqrt1->input(1) == real_div1 &&
|
||||
*add4->input(2) == *add2_y;
|
||||
return add5->input(kSecondIndex2) == mul4 && real_div4->input(kFirstIndex1) == real_div0 &&
|
||||
sqrt1->input(kFirstIndex1) == real_div1 && *add4->input(kSecondIndex2) == *add2_y;
|
||||
}
|
||||
|
||||
std::tuple<AnfNodePtr, AnfNodePtr> GetAdd0Add1Nodes(const AnfNodePtr &real_div0_anf, const AnfNodePtr &real_div1_anf) {
|
||||
|
|
@ -188,10 +193,9 @@ const AnfNodePtr LambNextMVWithDecayV1Rule::Process(const FuncGraphPtr &func_gra
|
|||
MS_LOG(EXCEPTION) << "create multiple outputs for fusion node fail!"
|
||||
<< " trace: " << trace::DumpSourceLines(node);
|
||||
}
|
||||
|
||||
(void)manager->Replace(add0, fusion_node_outputs[1]);
|
||||
(void)manager->Replace(add1, fusion_node_outputs[2]);
|
||||
(void)manager->Replace(add5, fusion_node_outputs[3]);
|
||||
(void)manager->Replace(add0, fusion_node_outputs[kFirstIndex1]);
|
||||
(void)manager->Replace(add1, fusion_node_outputs[kSecondIndex2]);
|
||||
(void)manager->Replace(add5, fusion_node_outputs[kThirdIndex3]);
|
||||
return fusion_node_outputs[0];
|
||||
}
|
||||
} // namespace opt
|
||||
|
|
|
|||
|
|
@ -25,16 +25,23 @@ namespace opt {
|
|||
const BaseRef LambUpdateWithLrV2::DefinePattern() const {
|
||||
const auto prim_greater = std::make_shared<Primitive>(kGreaterOpName);
|
||||
const auto prim_deal_div = std::make_shared<Primitive>(kRealDivOpName);
|
||||
const size_t kZeroIndex = 0;
|
||||
const size_t kFirstIndex = 1;
|
||||
const size_t kSecondIndex = 2;
|
||||
const size_t kThirdIndex = 3;
|
||||
const size_t kFourthIndex = 4;
|
||||
const size_t kFifthIndex = 5;
|
||||
const size_t kSixthIndex = 6;
|
||||
|
||||
VectorRef greater0({prim_greater, input_varptr_[0], input_varptr_[5]});
|
||||
VectorRef greater1({prim_greater, input_varptr_[1], input_varptr_[5]});
|
||||
VectorRef real_div0({prim_deal_div, input_varptr_[0], input_varptr_[1]});
|
||||
VectorRef select0({prim::kPrimSelect, greater1, real_div0, input_varptr_[6]});
|
||||
VectorRef select1({prim::kPrimSelect, greater0, select0, input_varptr_[6]});
|
||||
VectorRef mul0({prim::kPrimMul, select1, input_varptr_[2]});
|
||||
VectorRef mul1({prim::kPrimMul, mul0, input_varptr_[3]});
|
||||
VectorRef greater0({prim_greater, input_varptr_[kZeroIndex], input_varptr_[kFifthIndex]});
|
||||
VectorRef greater1({prim_greater, input_varptr_[kFirstIndex], input_varptr_[kFifthIndex]});
|
||||
VectorRef real_div0({prim_deal_div, input_varptr_[kZeroIndex], input_varptr_[kFirstIndex]});
|
||||
VectorRef select0({prim::kPrimSelect, greater1, real_div0, input_varptr_[kSixthIndex]});
|
||||
VectorRef select1({prim::kPrimSelect, greater0, select0, input_varptr_[kSixthIndex]});
|
||||
VectorRef mul0({prim::kPrimMul, select1, input_varptr_[kSecondIndex]});
|
||||
VectorRef mul1({prim::kPrimMul, mul0, input_varptr_[kThirdIndex]});
|
||||
|
||||
return VectorRef({prim::kPrimSub, input_varptr_[4], mul1});
|
||||
return VectorRef({prim::kPrimSub, input_varptr_[kFourthIndex], mul1});
|
||||
}
|
||||
|
||||
const AnfNodePtr LambUpdateWithLrV2::Process(const FuncGraphPtr &func_graph, const AnfNodePtr &node,
|
||||
|
|
|
|||
|
|
@ -53,9 +53,15 @@ const AnfNodePtr MomentumLossscaleFusion::Process(const FuncGraphPtr &func_graph
|
|||
MS_EXCEPTION_IF_NULL(func_graph);
|
||||
MS_EXCEPTION_IF_NULL(node);
|
||||
auto cnode = node->cast<CNodePtr>();
|
||||
constexpr size_t kFirstIndex = 1;
|
||||
constexpr size_t kSecondIndex = 2;
|
||||
constexpr size_t kThirdIndex = 3;
|
||||
constexpr size_t kFourthIndex = 4;
|
||||
constexpr size_t kFifthIndex = 5;
|
||||
constexpr size_t kSixthIndex = 6;
|
||||
MS_EXCEPTION_IF_NULL(cnode);
|
||||
CheckCNodeInputSize(cnode, kApplyMomentumInputTensorNum);
|
||||
AnfNodePtr mul = cnode->input(4);
|
||||
AnfNodePtr mul = cnode->input(kFourthIndex);
|
||||
MS_EXCEPTION_IF_NULL(mul);
|
||||
auto mul_cnode = mul->cast<CNodePtr>();
|
||||
MS_EXCEPTION_IF_NULL(mul_cnode);
|
||||
|
|
@ -74,13 +80,14 @@ const AnfNodePtr MomentumLossscaleFusion::Process(const FuncGraphPtr &func_graph
|
|||
}
|
||||
auto new_prim = std::make_shared<Primitive>(kFusedMulApplyMomentumOpName);
|
||||
auto depend_prim = NewValueNode(prim::kPrimDepend);
|
||||
auto depend = func_graph->NewCNode({depend_prim, cnode->input(5), cnode->input(6)}); // depend on monad
|
||||
depend->set_abstract(cnode->input(5)->abstract());
|
||||
depend->set_scope(cnode->input(5)->scope());
|
||||
auto depend =
|
||||
func_graph->NewCNode({depend_prim, cnode->input(kFifthIndex), cnode->input(kSixthIndex)}); // depend on monad
|
||||
depend->set_abstract(cnode->input(kFifthIndex)->abstract());
|
||||
depend->set_scope(cnode->input(kFifthIndex)->scope());
|
||||
std::vector<AnfNodePtr> new_node_inputs{NewValueNode(new_prim),
|
||||
cnode->input(1),
|
||||
cnode->input(2),
|
||||
cnode->input(3),
|
||||
cnode->input(kFirstIndex),
|
||||
cnode->input(kSecondIndex),
|
||||
cnode->input(kThirdIndex),
|
||||
mul_cnode->input(kMulInputTensorNum + 1 - value_node_index),
|
||||
depend,
|
||||
mul_cnode->input(value_node_index)};
|
||||
|
|
@ -88,7 +95,7 @@ const AnfNodePtr MomentumLossscaleFusion::Process(const FuncGraphPtr &func_graph
|
|||
MS_EXCEPTION_IF_NULL(new_node);
|
||||
AnfAlgo::CopyNodeAttrs(node, new_node);
|
||||
auto input_names_value = AnfAlgo::GetNodeAttr<std::vector<std::string>>(new_node, kAttrInputNames);
|
||||
input_names_value[3] = "x1";
|
||||
input_names_value[kThirdIndex] = "x1";
|
||||
input_names_value.emplace_back("x2");
|
||||
AnfAlgo::SetNodeAttr(kAttrInputNames, MakeValue(input_names_value), new_node);
|
||||
new_node->set_abstract(node->abstract());
|
||||
|
|
|
|||
|
|
@ -25,7 +25,7 @@
|
|||
namespace mindspore {
|
||||
namespace opt {
|
||||
namespace {
|
||||
bool GetMul(const FuncGraphPtr &graph, const CNodePtr &add, CNodePtr *mul, size_t *mul_index) {
|
||||
bool GetMul(const FuncGraphPtr &graph, const CNodePtr &add, CNodePtr *const mul, size_t *const mul_index) {
|
||||
MS_EXCEPTION_IF_NULL(graph);
|
||||
MS_EXCEPTION_IF_NULL(add);
|
||||
|
||||
|
|
|
|||
|
|
@ -28,10 +28,13 @@ bool CheckShapeDimInfo(const std::vector<size_t> &shape) {
|
|||
if (shape.empty()) {
|
||||
return false;
|
||||
}
|
||||
if (shape.size() == 1 && shape[0] % kCubeSize != 0) {
|
||||
constexpr auto kShapeSize1 = 1;
|
||||
constexpr auto kShapeSize2 = 2;
|
||||
if (shape.size() == kShapeSize1 && shape[0] % kCubeSize != 0) {
|
||||
return false;
|
||||
}
|
||||
return !(shape.size() >= 2 && (shape[shape.size() - 1] % kCubeSize != 0 || shape[shape.size() - 2] % kCubeSize != 0));
|
||||
return !(shape.size() >= kShapeSize2 &&
|
||||
(shape[shape.size() - 1] % kCubeSize != 0 || shape[shape.size() - kShapeSize2] % kCubeSize != 0));
|
||||
}
|
||||
} // namespace
|
||||
|
||||
|
|
|
|||
Loading…
Reference in New Issue