From: @zhupuxu
Reviewed-by: @jjfeing,@kisnwang
Signed-off-by: @kisnwang
This commit is contained in:
mindspore-ci-bot 2021-05-26 09:14:39 +08:00 committed by Gitee
commit 22d62f20ff
11 changed files with 128 additions and 77 deletions

View File

@ -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

View File

@ -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());

View File

@ -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};

View File

@ -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(

View File

@ -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];
}

View File

@ -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;
}

View File

@ -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

View File

@ -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,

View File

@ -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());

View File

@ -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);

View File

@ -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