forked from huawei/mindspore2022
update resnet network performence.
less_bn pattern update. fix clang-format !19495 add allreduce operators by parameters for lessBN and add group parameters generator for lessBN, gc and grad freeze Merge pull request !19495 from jinjiali-kali/less_bn update resnet acc_mode scripts.
This commit is contained in:
parent
0305441854
commit
9975c6a3a8
|
|
@ -177,7 +177,8 @@ OptimizeIRPassLib::OptimizeIRPassLib() {
|
|||
|
||||
// Accelerated Algorithm
|
||||
less_batch_normalization_ =
|
||||
MakeSubstitution(std::make_shared<LessBatchNormalization>(), "less_batch_normalization", prim::kPrimAdd);
|
||||
MakeSubstitution(std::make_shared<LessBatchNormalization>(), "less_batch_normalization",
|
||||
{prim::kPrimAdd, prim::kPrimRelu6, prim::kPrimMatMul, prim::kPrimMakeTuple, prim::kPrimMaxPool});
|
||||
|
||||
// inline
|
||||
inline_ = MakeSubstitution(std::make_shared<Inliner>(), "inline", IsCNodeGraph);
|
||||
|
|
|
|||
|
|
@ -31,8 +31,8 @@ constexpr auto kFirstBranchPattern1 = 12;
|
|||
constexpr auto kSecondBranchPattern1 = 3;
|
||||
constexpr auto kFirstBranchStartIndexPattern1 = 4;
|
||||
constexpr auto kFirstBranchEndIndexPattern1 = 11;
|
||||
constexpr auto kSecondBranchStartIndexPattern1 = 12;
|
||||
constexpr auto kSecondBranchEndIndexPattern1 = 14;
|
||||
constexpr auto kSecondBranchStartIndexPattern1 = kFirstBranchPattern1;
|
||||
constexpr auto kSecondBranchEndIndexPattern1 = 2 + kFirstBranchPattern1;
|
||||
const std::vector<kStructureTuple> ResidualStructureBasePattern{
|
||||
{kFirstBranchPattern1,
|
||||
{prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimConv2D, prim::kPrimRelu},
|
||||
|
|
@ -47,8 +47,8 @@ constexpr auto kFirstBranchPattern2 = 12;
|
|||
constexpr auto kSecondBranchPattern2 = 1;
|
||||
constexpr auto kFirstBranchStartIndexPattern2 = 4;
|
||||
constexpr auto kFirstBranchEndIndexPattern2 = 11;
|
||||
constexpr auto kSecondBranchStartIndexPattern2 = 12;
|
||||
constexpr auto kSecondBranchEndIndexPattern2 = 13;
|
||||
constexpr auto kSecondBranchStartIndexPattern2 = kFirstBranchPattern2;
|
||||
constexpr auto kSecondBranchEndIndexPattern2 = 1 + kSecondBranchPattern2;
|
||||
const std::vector<kStructureTuple> ResidualStructureShortCutPattern{
|
||||
{kFirstBranchPattern2,
|
||||
{prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimConv2D, prim::kPrimRelu},
|
||||
|
|
@ -61,8 +61,8 @@ constexpr auto kFirstBranchPattern3 = 11;
|
|||
constexpr auto kSecondBranchPattern3 = 3;
|
||||
constexpr auto kFirstBranchStartIndexPattern3 = 4;
|
||||
constexpr auto kFirstBranchEndIndexPattern3 = 10;
|
||||
constexpr auto kSecondBranchStartIndexPattern3 = 11;
|
||||
constexpr auto kSecondBranchEndIndexPattern3 = 13;
|
||||
constexpr auto kSecondBranchStartIndexPattern3 = kFirstBranchPattern3;
|
||||
constexpr auto kSecondBranchEndIndexPattern3 = 2 + kFirstBranchPattern3;
|
||||
const std::vector<kStructureTuple> ResidualStructureFirstStepPattern{
|
||||
{kFirstBranchPattern3,
|
||||
{prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimConv2D, prim::kPrimRelu, prim::kPrimTupleGetItem,
|
||||
|
|
@ -73,15 +73,13 @@ const std::vector<kStructureTuple> ResidualStructureFirstStepPattern{
|
|||
{prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimConv2D},
|
||||
{kSecondBranchStartIndexPattern3, kSecondBranchEndIndexPattern3}}};
|
||||
// Pattern 4
|
||||
// Add -> BatchNorm -> Conv2D -> Relu ... -> End
|
||||
// ↘ BatchNorm -> Conv2D -> -> -> -> ↗
|
||||
constexpr auto kFirstBranchPattern4 = 8;
|
||||
constexpr auto kSecondBranchPattern4 = 3;
|
||||
constexpr auto kFirstBranchStartIndexPattern4 = 4;
|
||||
constexpr auto kFirstBranchEndIndexPattern4 = 6;
|
||||
constexpr auto kSecondBranchStartIndexPattern4 = 8;
|
||||
constexpr auto kSecondBranchEndIndexPattern4 = 11;
|
||||
const std::vector<kStructureTuple> BasicStructureBasePattern{
|
||||
constexpr auto kSecondBranchStartIndexPattern4 = kFirstBranchPattern4;
|
||||
constexpr auto kSecondBranchEndIndexPattern4 = 3 + kFirstBranchPattern4;
|
||||
const std::vector<kStructureTuple> BasicStructBasePattern{
|
||||
{kFirstBranchPattern4,
|
||||
{prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimConv2D, prim::kPrimRelu},
|
||||
{kFirstBranchStartIndexPattern4, kFirstBranchEndIndexPattern4}},
|
||||
|
|
@ -89,37 +87,163 @@ const std::vector<kStructureTuple> BasicStructureBasePattern{
|
|||
{prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimConv2D},
|
||||
{kSecondBranchStartIndexPattern4, kSecondBranchEndIndexPattern4}}};
|
||||
// Pattern 5
|
||||
// Add -> BatchNorm -> Conv2D -> Relu ... -> End
|
||||
// ↘ -> -> -> -> Relu -> -> -> -> ↗
|
||||
constexpr auto kFirstBranchPattern5 = 8;
|
||||
constexpr auto kFirstBranchPattern5 = 7;
|
||||
constexpr auto kSecondBranchPattern5 = 1;
|
||||
constexpr auto kFirstBranchStartIndexPattern5 = 4;
|
||||
constexpr auto kFirstBranchEndIndexPattern5 = 6;
|
||||
constexpr auto kSecondBranchStartIndexPattern5 = 8;
|
||||
constexpr auto kSecondBranchEndIndexPattern5 = 11;
|
||||
const std::vector<kStructureTuple> BasicStructureShortCutPattern{
|
||||
constexpr auto kSecondBranchStartIndexPattern5 = kFirstBranchPattern5;
|
||||
constexpr auto kSecondBranchEndIndexPattern5 = 3 + kFirstBranchPattern5;
|
||||
const std::vector<kStructureTuple> BasicStructFirstStepPattern{
|
||||
{kFirstBranchPattern5,
|
||||
{prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimConv2D, prim::kPrimRelu},
|
||||
{prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimConv2D, prim::kPrimRelu, prim::kPrimTupleGetItem,
|
||||
prim::kPrimBatchNorm, prim::kPrimConv2D},
|
||||
{kFirstBranchStartIndexPattern5, kFirstBranchEndIndexPattern5}},
|
||||
{kSecondBranchPattern5, {prim::kPrimRelu}, {kSecondBranchStartIndexPattern5, kSecondBranchEndIndexPattern5}}};
|
||||
{kSecondBranchPattern5, {prim::kPrimMaxPool}, {kSecondBranchStartIndexPattern5, kSecondBranchEndIndexPattern5}}};
|
||||
// Pattern 6
|
||||
// Add -> BatchNorm -> Conv2D -> Relu ... -> End
|
||||
// ↘ -> -> -> -> MaxPool -> -> -> ↗
|
||||
constexpr auto kFirstBranchPattern6 = 7;
|
||||
constexpr auto kFirstBranchPattern6 = 8;
|
||||
constexpr auto kSecondBranchPattern6 = 1;
|
||||
constexpr auto kFirstBranchStartIndexPattern6 = 4;
|
||||
constexpr auto kFirstBranchEndIndexPattern6 = 6;
|
||||
constexpr auto kSecondBranchStartIndexPattern6 = 7;
|
||||
constexpr auto kSecondBranchEndIndexPattern6 = 10;
|
||||
const std::vector<kStructureTuple> BasicStructureFirstStepPattern{
|
||||
constexpr auto kSecondBranchStartIndexPattern6 = kFirstBranchPattern6;
|
||||
constexpr auto kSecondBranchEndIndexPattern6 = 3 + kFirstBranchPattern6;
|
||||
const std::vector<kStructureTuple> BasicStructShortCutPattern{
|
||||
{kFirstBranchPattern6,
|
||||
{prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimConv2D, prim::kPrimRelu, prim::kPrimTupleGetItem,
|
||||
prim::kPrimBatchNorm, prim::kPrimConv2D},
|
||||
{prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimConv2D, prim::kPrimRelu},
|
||||
{kFirstBranchStartIndexPattern6, kFirstBranchEndIndexPattern6}},
|
||||
{kSecondBranchPattern6, {prim::kPrimMaxPool}, {kSecondBranchStartIndexPattern6, kSecondBranchEndIndexPattern6}}};
|
||||
static const std::vector<std::vector<kStructureTuple>> kNeedMatchPattern = {
|
||||
ResidualStructureBasePattern, ResidualStructureShortCutPattern, ResidualStructureFirstStepPattern,
|
||||
BasicStructureBasePattern, BasicStructureShortCutPattern, BasicStructureFirstStepPattern};
|
||||
{kSecondBranchPattern6, {prim::kPrimRelu}, {kSecondBranchStartIndexPattern6, kSecondBranchEndIndexPattern6}}};
|
||||
// Pattern 7
|
||||
constexpr auto kFirstBranchPattern7 = 1;
|
||||
constexpr auto kSecondBranchPattern7 = 13;
|
||||
constexpr auto kFirstBranchStartIndexPattern7 = SIZE_MAX;
|
||||
constexpr auto kFirstBranchEndIndexPattern7 = SIZE_MAX;
|
||||
constexpr auto kSecondBranchStartIndexPattern7 = 7;
|
||||
constexpr auto kSecondBranchEndIndexPattern7 = 10;
|
||||
const std::vector<kStructureTuple> InvertedResidualShortCutPattern{
|
||||
{kFirstBranchPattern7,
|
||||
{prim::kPrimTupleGetItem, prim::kPrimBatchNorm},
|
||||
{kFirstBranchStartIndexPattern7, kFirstBranchEndIndexPattern7}},
|
||||
{kSecondBranchPattern7,
|
||||
{prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimConv2D, prim::kPrimRelu6, prim::kPrimTupleGetItem,
|
||||
prim::kPrimBatchNorm, prim::kPrimConv2D, prim::kPrimRelu6, prim::kPrimTupleGetItem, prim::kPrimBatchNorm,
|
||||
prim::kPrimConv2D, prim::kPrimTupleGetItem, prim::kPrimBatchNorm},
|
||||
{kSecondBranchStartIndexPattern7, kSecondBranchEndIndexPattern7}}};
|
||||
// Pattern 8
|
||||
constexpr auto kFirstBranchPattern8 = 4;
|
||||
constexpr auto kFirstBranchStartIndexPattern8 = 0;
|
||||
constexpr auto kFirstBranchEndIndexPattern8 = 3;
|
||||
const std::vector<kStructureTuple> InvertedResidualPattern{
|
||||
{kFirstBranchPattern8,
|
||||
{prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimConv2D, prim::kPrimAdd},
|
||||
{kFirstBranchStartIndexPattern8, kFirstBranchEndIndexPattern8}}};
|
||||
// Pattern 9
|
||||
constexpr auto kFirstBranchPattern9 = 1;
|
||||
constexpr auto kSecondBranchPattern9 = 12;
|
||||
constexpr auto kFirstBranchStartIndexPattern9 = SIZE_MAX;
|
||||
constexpr auto kFirstBranchEndIndexPattern9 = SIZE_MAX;
|
||||
constexpr auto kSecondBranchStartIndexPattern9 = 7;
|
||||
constexpr auto kSecondBranchEndIndexPattern9 = 10;
|
||||
const std::vector<kStructureTuple> InvertedResidualShortCutPattern2{
|
||||
{kFirstBranchPattern9, {prim::kPrimAdd}, {kFirstBranchStartIndexPattern9, kFirstBranchEndIndexPattern9}},
|
||||
{kSecondBranchPattern9,
|
||||
{prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimConv2D, prim::kPrimRelu6, prim::kPrimTupleGetItem,
|
||||
prim::kPrimBatchNorm, prim::kPrimConv2D, prim::kPrimRelu6, prim::kPrimTupleGetItem, prim::kPrimBatchNorm,
|
||||
prim::kPrimConv2D, prim::kPrimAdd},
|
||||
{kSecondBranchStartIndexPattern9, kSecondBranchEndIndexPattern9}}};
|
||||
// Pattern 10
|
||||
constexpr auto kFirstBranchPattern10 = 5;
|
||||
constexpr auto kFirstBranchStartIndexPattern10 = 0;
|
||||
constexpr auto kFirstBranchEndIndexPattern10 = 4;
|
||||
const std::vector<kStructureTuple> InvertedResidualPattern2{
|
||||
{kFirstBranchPattern10,
|
||||
{prim::kPrimReduceMean, prim::kPrimRelu6, prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimConv2D},
|
||||
{kFirstBranchStartIndexPattern10, kFirstBranchEndIndexPattern10}}};
|
||||
// Pattern 11
|
||||
constexpr auto kFirstBranchPattern11 = 17;
|
||||
constexpr auto kFirstBranchStartIndexPattern11 = 3;
|
||||
constexpr auto kFirstBranchEndIndexPattern11 = 6;
|
||||
const std::vector<kStructureTuple> InvertedResidualPattern3{
|
||||
{kFirstBranchPattern11,
|
||||
{prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimConv2D, prim::kPrimRelu6, prim::kPrimTupleGetItem,
|
||||
prim::kPrimBatchNorm, prim::kPrimConv2D, prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimConv2D,
|
||||
prim::kPrimRelu6, prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimConv2D, prim::kPrimRelu6,
|
||||
prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimConv2D},
|
||||
{kFirstBranchStartIndexPattern11, kFirstBranchEndIndexPattern11}}};
|
||||
// Pattern 12
|
||||
constexpr auto kFirstBranchPattern12 = 1;
|
||||
constexpr auto kSecondBranchPattern12 = 9;
|
||||
constexpr auto kFirstBranchStartIndexPattern12 = SIZE_MAX;
|
||||
constexpr auto kFirstBranchEndIndexPattern12 = SIZE_MAX;
|
||||
constexpr auto kSecondBranchStartIndexPattern12 = kFirstBranchPattern12 + 5;
|
||||
constexpr auto kSecondBranchEndIndexPattern12 = kFirstBranchPattern12 + 8;
|
||||
const std::vector<kStructureTuple> DenseBlockShortCutPattern{
|
||||
{kFirstBranchPattern12, {prim::kPrimConcat}, {kFirstBranchStartIndexPattern12, kFirstBranchEndIndexPattern12}},
|
||||
{kSecondBranchPattern12,
|
||||
{prim::kPrimConv2D, prim::kPrimRelu, prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimConv2D,
|
||||
prim::kPrimRelu, prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimConcat},
|
||||
{kSecondBranchStartIndexPattern12, kSecondBranchEndIndexPattern12}}};
|
||||
// Pattern 13
|
||||
constexpr auto kFirstBranchPattern13 = 5;
|
||||
constexpr auto kFirstBranchStartIndexPattern13 = 0;
|
||||
constexpr auto kFirstBranchEndIndexPattern13 = 4;
|
||||
const std::vector<kStructureTuple> DenseBlockPattern{
|
||||
{kFirstBranchPattern13,
|
||||
{prim::kPrimConv2D, prim::kPrimRelu, prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimConcat},
|
||||
{kFirstBranchStartIndexPattern13, kFirstBranchEndIndexPattern13}}};
|
||||
// Pattern 14
|
||||
constexpr auto kFirstBranchPattern14 = 9;
|
||||
constexpr auto kSecondBranchPattern14 = 1;
|
||||
constexpr auto kFirstBranchStartIndexPattern14 = 5;
|
||||
constexpr auto kFirstBranchEndIndexPattern14 = 8;
|
||||
constexpr auto kSecondBranchStartIndexPattern14 = SIZE_MAX;
|
||||
constexpr auto kSecondBranchEndIndexPattern14 = SIZE_MAX;
|
||||
const std::vector<kStructureTuple> DenseBlockShortCutPattern2{
|
||||
{kFirstBranchPattern14,
|
||||
{prim::kPrimConv2D, prim::kPrimRelu, prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimConv2D,
|
||||
prim::kPrimRelu, prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimConcat},
|
||||
{kFirstBranchStartIndexPattern14, kFirstBranchEndIndexPattern14}},
|
||||
{kSecondBranchPattern14, {prim::kPrimConcat}, {kSecondBranchStartIndexPattern14, kSecondBranchEndIndexPattern14}}};
|
||||
// Pattern 15
|
||||
constexpr auto kFirstBranchPattern15 = 9;
|
||||
constexpr auto kSecondBranchPattern15 = 1;
|
||||
constexpr auto kFirstBranchStartIndexPattern15 = 0;
|
||||
constexpr auto kFirstBranchEndIndexPattern15 = 4;
|
||||
constexpr auto kSecondBranchStartIndexPattern15 = SIZE_MAX;
|
||||
constexpr auto kSecondBranchEndIndexPattern15 = SIZE_MAX;
|
||||
const std::vector<kStructureTuple> DenseBlockPoolPattern{
|
||||
{kFirstBranchPattern15,
|
||||
{prim::kPrimConv2D, prim::kPrimRelu, prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimConv2D,
|
||||
prim::kPrimRelu, prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimMaxPool},
|
||||
{kFirstBranchStartIndexPattern15, kFirstBranchEndIndexPattern15}},
|
||||
{kSecondBranchPattern15, {prim::kPrimConcat}, {kSecondBranchStartIndexPattern15, kSecondBranchEndIndexPattern15}}};
|
||||
// Pattern 16
|
||||
constexpr auto kFirstBranchPattern16 = 1;
|
||||
constexpr auto kSecondBranchPattern16 = 9;
|
||||
constexpr auto kFirstBranchStartIndexPattern16 = SIZE_MAX;
|
||||
constexpr auto kFirstBranchEndIndexPattern16 = SIZE_MAX;
|
||||
constexpr auto kSecondBranchStartIndexPattern16 = kFirstBranchPattern16;
|
||||
constexpr auto kSecondBranchEndIndexPattern16 = kFirstBranchPattern16 + 4;
|
||||
const std::vector<kStructureTuple> DenseBlockPoolPatter2{
|
||||
{kFirstBranchPattern16, {prim::kPrimConcat}, {kFirstBranchStartIndexPattern16, kFirstBranchEndIndexPattern16}},
|
||||
{kSecondBranchPattern16,
|
||||
{prim::kPrimConv2D, prim::kPrimRelu, prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimConv2D,
|
||||
prim::kPrimRelu, prim::kPrimTupleGetItem, prim::kPrimBatchNorm, prim::kPrimMaxPool},
|
||||
{kSecondBranchStartIndexPattern16, kSecondBranchEndIndexPattern16}}};
|
||||
static const std::vector<std::vector<kStructureTuple>> kNeedMatchPattern = {ResidualStructureBasePattern,
|
||||
ResidualStructureShortCutPattern,
|
||||
ResidualStructureFirstStepPattern,
|
||||
BasicStructBasePattern,
|
||||
BasicStructFirstStepPattern,
|
||||
BasicStructShortCutPattern,
|
||||
InvertedResidualShortCutPattern,
|
||||
InvertedResidualPattern,
|
||||
InvertedResidualShortCutPattern2,
|
||||
InvertedResidualPattern2,
|
||||
InvertedResidualPattern3,
|
||||
DenseBlockShortCutPattern,
|
||||
DenseBlockPattern,
|
||||
DenseBlockShortCutPattern2,
|
||||
DenseBlockPoolPattern,
|
||||
DenseBlockPoolPatter2};
|
||||
const std::set<PrimitivePtr> kNeedRemoveNodeSet{
|
||||
prim::kPrimLoad, prim::kPrimRefToEmbed, prim::kPrimApplyMomentum, prim::kPrimMomentum,
|
||||
prim::kPrimApplyFtrl, prim::kPrimSGD, prim::kPrimApplyRMSProp, prim::kPrimAdam};
|
||||
|
|
@ -286,7 +410,13 @@ AnfNodePtr LessBatchNormalization::operator()(const OptimizerPtr &optimizer, con
|
|||
sum_match_node += std::get<0>(t);
|
||||
total_match_node_.emplace_back(sum_match_node);
|
||||
});
|
||||
AnfVisitor::Match(prim::kPrimAdd, {IsCNode, IsCNode})(node);
|
||||
auto cnode = node->cast<CNodePtr>();
|
||||
if (cnode == nullptr || cnode->inputs().empty()) {
|
||||
return nullptr;
|
||||
}
|
||||
auto prim = GetValueNode<PrimitivePtr>(cnode->input(0));
|
||||
std::vector<PredicateFuncType> funcs(cnode->inputs().size() - 1, IsCNode);
|
||||
AnfVisitor::Match(prim, funcs)(node);
|
||||
if (is_match_) {
|
||||
break;
|
||||
}
|
||||
|
|
|
|||
|
|
@ -40,6 +40,7 @@ class AutoAcc:
|
|||
def __init__(self, level, kwargs):
|
||||
if level not in _acc_config_level.keys():
|
||||
level = 'O0'
|
||||
self.level = level
|
||||
acc_config = _acc_config_level[level]
|
||||
self._acc_config = acc_config
|
||||
self._fn_flag = True
|
||||
|
|
@ -66,12 +67,11 @@ class AutoAcc:
|
|||
optimizer_process = OptimizerProcess(optimizer)
|
||||
group_params = self._param_processer.assign_parameter_group(network.trainable_params(),
|
||||
self._gradient_groups)
|
||||
optimizer_process.origin_params = self._param_processer.generate_group_params(group_params)
|
||||
optimizer_process.origin_params = \
|
||||
self._param_processer.generate_group_params(group_params, optimizer_process.origin_params)
|
||||
if self._gc_flag:
|
||||
parameters = optimizer_process.add_grad_centralization()
|
||||
else:
|
||||
parameters = optimizer_process.origin_params
|
||||
optimizer = optimizer_process.generate_new_optimizer(parameters)
|
||||
optimizer_process.add_grad_centralization()
|
||||
optimizer = optimizer_process.generate_new_optimizer()
|
||||
|
||||
if self._acc_config["grad_freeze"]:
|
||||
freeze_processer = GradientFreeze(self._param_groups, self._freeze_type,
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@
|
|||
# limitations under the License.
|
||||
# ============================================================================
|
||||
"""base process"""
|
||||
import copy
|
||||
from mindspore.nn.cell import Cell
|
||||
from mindspore.nn.optim import LARS
|
||||
from mindspore import log as logger
|
||||
|
|
@ -58,24 +59,37 @@ class OptimizerProcess:
|
|||
|
||||
if isinstance(parameters[0], Parameter):
|
||||
logger.warning("Only group parameters support gradient centralization.")
|
||||
return parameters
|
||||
return
|
||||
|
||||
change_dict = parameters[0]
|
||||
if 'order_params' in change_dict.keys():
|
||||
logger.warning("Only support normal parameters for gradient centralization.")
|
||||
return parameters
|
||||
group_params = []
|
||||
for group_param in parameters:
|
||||
if 'order_params' in group_param.keys():
|
||||
group_params.append(group_param)
|
||||
continue
|
||||
params_gc_value = []
|
||||
params_value = []
|
||||
for param in group_param['params']:
|
||||
if 'beta' not in param.name and 'gamma' not in param.name and 'bias' not in param.name:
|
||||
params_gc_value.append(param)
|
||||
else:
|
||||
params_value.append(param)
|
||||
if params_gc_value:
|
||||
new_group_param = copy.deepcopy(group_param)
|
||||
new_group_param['params'] = params_gc_value
|
||||
new_group_param['grad_centralization'] = True
|
||||
group_params.append(new_group_param)
|
||||
if params_value:
|
||||
new_group_param = copy.deepcopy(group_param)
|
||||
new_group_param['params'] = params_value
|
||||
group_params.append(new_group_param)
|
||||
self.origin_params = group_params
|
||||
|
||||
change_dict['grad_centralization'] = True
|
||||
self.origin_params[0] = change_dict
|
||||
|
||||
return self.origin_params
|
||||
|
||||
def generate_new_optimizer(self, params):
|
||||
def generate_new_optimizer(self):
|
||||
"""Generate new optimizer."""
|
||||
if not self.is_lars:
|
||||
opt = self.opt_class(params=params, **self.opt_init_args)
|
||||
opt = self.opt_class(params=self.origin_params, **self.opt_init_args)
|
||||
else:
|
||||
opt = LARS(self.opt_class(params=params, **self.opt_init_args), **self.lars_init_args)
|
||||
opt = LARS(self.opt_class(params=self.origin_params, **self.opt_init_args), **self.lars_init_args)
|
||||
|
||||
return opt
|
||||
|
||||
|
|
@ -103,19 +117,46 @@ class ParameterProcess:
|
|||
parameters[i].comm_fusion = self._parameter_indices
|
||||
return parameters
|
||||
|
||||
def generate_group_params(self, parameters):
|
||||
def generate_group_params(self, parameters, origin_params):
|
||||
"""Generate group parameters."""
|
||||
decayed_params = []
|
||||
no_decayed_params = []
|
||||
for param in parameters:
|
||||
if 'beta' not in param.name and 'gamma' not in param.name and 'bias' not in param.name:
|
||||
decayed_params.append(param)
|
||||
else:
|
||||
no_decayed_params.append(param)
|
||||
group_params = [{'params': decayed_params, 'weight_decay': 0.0001},
|
||||
{'params': no_decayed_params},
|
||||
{'order_params': parameters}]
|
||||
origin_params_copy = origin_params
|
||||
if origin_params_copy is not None:
|
||||
if not isinstance(origin_params_copy, list):
|
||||
origin_params_copy = list(origin_params_copy)
|
||||
|
||||
if not origin_params_copy:
|
||||
raise ValueError("Optimizer got an empty parameter list.")
|
||||
|
||||
if not isinstance(origin_params_copy[0], (dict, Parameter)):
|
||||
raise TypeError("Only a list of Parameter or dict can be supported.")
|
||||
|
||||
if isinstance(origin_params_copy[0], Parameter):
|
||||
group_params = [{"params": parameters}]
|
||||
else:
|
||||
group_params = []
|
||||
params_name = [param.name for param in parameters]
|
||||
new_params_count = copy.deepcopy(params_name)
|
||||
for group_param in origin_params_copy:
|
||||
if 'order_params' in group_param.keys():
|
||||
new_group_param = copy.deepcopy(group_param)
|
||||
new_group_param['order_params'] = parameters
|
||||
group_params.append(new_group_param)
|
||||
continue
|
||||
params_value = []
|
||||
for param in group_param['params']:
|
||||
if param.name in params_name:
|
||||
index = params_name.index(param.name)
|
||||
params_value.append(parameters[index])
|
||||
new_params_count.remove(param.name)
|
||||
new_group_param = copy.deepcopy(group_param)
|
||||
new_group_param['params'] = params_value
|
||||
group_params.append(new_group_param)
|
||||
if new_params_count:
|
||||
params_value = []
|
||||
for param in new_params_count:
|
||||
index = params_name.index(param)
|
||||
params_value.append(parameters[index])
|
||||
group_params.append({"params": params_value})
|
||||
return group_params
|
||||
|
||||
_gradient_accumulation_op = C.MultitypeFuncGraph("gradient_accumulation_op")
|
||||
|
|
|
|||
|
|
@ -235,7 +235,8 @@ class GradientFreeze:
|
|||
train_para_groups = self.split_parameters_groups(
|
||||
network, self._param_groups)
|
||||
for i in range(self._param_groups):
|
||||
train_para_groups[i] = self._param_processer.generate_group_params(train_para_groups[i])
|
||||
train_para_groups[i] = self._param_processer.generate_group_params(train_para_groups[i],
|
||||
optimizer.init_params['params'])
|
||||
train_strategy = self.generate_freeze_index_sequence(
|
||||
self._param_groups, self._freeze_type, self._freeze_p, self._total_steps)
|
||||
optimizer = FreezeOpt(optimizer, train_para_groups, train_strategy)
|
||||
|
|
@ -248,7 +249,7 @@ def freeze_cell(reducer_flag, network, optimizer, sens, grad, use_grad_accumulat
|
|||
if reducer_flag:
|
||||
param_processer = ParameterProcess()
|
||||
grad_reducers = (DistributedGradReducer(param_processer.assign_parameter_group(opt.parameters),
|
||||
mean, degree, param_fusion=True) for opt in optimizer.opts)
|
||||
mean, degree) for opt in optimizer.opts)
|
||||
freeze_nets = tuple(_TrainFreezeCell(network, sens, grad, reducer,
|
||||
use_grad_accumulation, opt, max_accumulation_step)
|
||||
for reducer, opt in zip(grad_reducers, optimizer.opts))
|
||||
|
|
|
|||
|
|
@ -16,6 +16,7 @@
|
|||
from mindspore.ops import functional as F, composite as C, operations as P
|
||||
from mindspore._checkparam import Validator as validator
|
||||
from .optimizer import Optimizer
|
||||
from .optimizer import opt_init_args_register
|
||||
|
||||
_ada_grad_opt = C.MultitypeFuncGraph("ada_grad_opt")
|
||||
|
||||
|
|
@ -144,6 +145,7 @@ class Adagrad(Optimizer):
|
|||
>>> model = Model(net, loss_fn=loss, optimizer=optim)
|
||||
"""
|
||||
|
||||
@opt_init_args_register
|
||||
def __init__(self, params, accum=0.1, learning_rate=0.001,
|
||||
update_slots=True, loss_scale=1.0, weight_decay=0.0):
|
||||
super(Adagrad, self).__init__(learning_rate, params, weight_decay, loss_scale)
|
||||
|
|
|
|||
|
|
@ -25,6 +25,7 @@ from mindspore.common.tensor import Tensor
|
|||
from mindspore._checkparam import Validator as validator
|
||||
from mindspore._checkparam import Rel
|
||||
from .optimizer import Optimizer
|
||||
from .optimizer import opt_init_args_register
|
||||
|
||||
_adam_opt = C.MultitypeFuncGraph("adam_opt")
|
||||
_scaler_one = Tensor(1, mstype.int32)
|
||||
|
|
@ -311,6 +312,7 @@ class Adam(Optimizer):
|
|||
>>> model = Model(net, loss_fn=loss, optimizer=optim)
|
||||
"""
|
||||
|
||||
@opt_init_args_register
|
||||
def __init__(self, params, learning_rate=1e-3, beta1=0.9, beta2=0.999, eps=1e-8, use_locking=False,
|
||||
use_nesterov=False, weight_decay=0.0, loss_scale=1.0):
|
||||
super(Adam, self).__init__(learning_rate, params, weight_decay, loss_scale)
|
||||
|
|
|
|||
|
|
@ -19,6 +19,7 @@ import mindspore.common.dtype as mstype
|
|||
from mindspore._checkparam import Validator as validator
|
||||
from mindspore._checkparam import Rel
|
||||
from .optimizer import Optimizer, _apply_decay, _grad_scale
|
||||
from .optimizer import opt_init_args_register
|
||||
|
||||
_ftrl_opt = C.MultitypeFuncGraph("ftrl_opt")
|
||||
|
||||
|
|
@ -191,6 +192,8 @@ class FTRL(Optimizer):
|
|||
>>> loss = nn.SoftmaxCrossEntropyWithLogits()
|
||||
>>> model = Model(net, loss_fn=loss, optimizer=optim)
|
||||
"""
|
||||
|
||||
@opt_init_args_register
|
||||
def __init__(self, params, initial_accum=0.1, learning_rate=0.001, lr_power=-0.5, l1=0.0, l2=0.0,
|
||||
use_locking=False, loss_scale=1.0, weight_decay=0.0):
|
||||
super(FTRL, self).__init__(learning_rate, params, weight_decay, loss_scale=loss_scale)
|
||||
|
|
|
|||
|
|
@ -25,6 +25,7 @@ from mindspore.common.tensor import Tensor
|
|||
from mindspore._checkparam import Validator as validator
|
||||
from mindspore._checkparam import Rel
|
||||
from .optimizer import Optimizer
|
||||
from .optimizer import opt_init_args_register
|
||||
from .. import layer
|
||||
|
||||
|
||||
|
|
@ -266,6 +267,7 @@ class Lamb(Optimizer):
|
|||
>>> model = Model(net, loss_fn=loss, optimizer=optim)
|
||||
"""
|
||||
|
||||
@opt_init_args_register
|
||||
def __init__(self, params, learning_rate, beta1=0.9, beta2=0.999, eps=1e-6, weight_decay=0.0):
|
||||
super(Lamb, self).__init__(learning_rate, params, weight_decay)
|
||||
_check_param_value(beta1, beta2, eps, self.cls_name)
|
||||
|
|
|
|||
|
|
@ -23,6 +23,7 @@ from mindspore.common.tensor import Tensor
|
|||
from mindspore._checkparam import Validator as validator
|
||||
from mindspore._checkparam import Rel
|
||||
from .optimizer import Optimizer
|
||||
from .optimizer import opt_init_args_register
|
||||
|
||||
_lazy_adam_opt = C.MultitypeFuncGraph("lazy_adam_opt")
|
||||
|
||||
|
|
@ -231,6 +232,7 @@ class LazyAdam(Optimizer):
|
|||
>>> model = Model(net, loss_fn=loss, optimizer=optim)
|
||||
"""
|
||||
|
||||
@opt_init_args_register
|
||||
def __init__(self, params, learning_rate=1e-3, beta1=0.9, beta2=0.999, eps=1e-8, use_locking=False,
|
||||
use_nesterov=False, weight_decay=0.0, loss_scale=1.0):
|
||||
super(LazyAdam, self).__init__(learning_rate, params, weight_decay, loss_scale)
|
||||
|
|
|
|||
|
|
@ -18,6 +18,7 @@ from mindspore.common import Tensor
|
|||
import mindspore.common.dtype as mstype
|
||||
from mindspore._checkparam import Validator as validator
|
||||
from .optimizer import Optimizer
|
||||
from .optimizer import opt_init_args_register
|
||||
|
||||
_proximal_ada_grad_opt = C.MultitypeFuncGraph("proximal_ada_grad_opt")
|
||||
|
||||
|
|
@ -167,6 +168,7 @@ class ProximalAdagrad(Optimizer):
|
|||
>>> model = Model(net, loss_fn=loss, optimizer=optim)
|
||||
"""
|
||||
|
||||
@opt_init_args_register
|
||||
def __init__(self, params, accum=0.1, learning_rate=0.001, l1=0.0, l2=0.0,
|
||||
use_locking=False, loss_scale=1.0, weight_decay=0.0):
|
||||
super(ProximalAdagrad, self).__init__(learning_rate, params, weight_decay, loss_scale)
|
||||
|
|
|
|||
|
|
@ -16,6 +16,7 @@
|
|||
from mindspore.ops import functional as F, composite as C, operations as P
|
||||
from mindspore._checkparam import Validator as validator
|
||||
from .optimizer import Optimizer
|
||||
from .optimizer import opt_init_args_register
|
||||
|
||||
_rmsprop_opt = C.MultitypeFuncGraph("rmsprop_opt")
|
||||
_centered_rmsprop_opt = C.MultitypeFuncGraph("rmsprop_opt")
|
||||
|
|
@ -175,6 +176,8 @@ class RMSProp(Optimizer):
|
|||
>>> loss = nn.SoftmaxCrossEntropyWithLogits()
|
||||
>>> model = Model(net, loss_fn=loss, optimizer=optim)
|
||||
"""
|
||||
|
||||
@opt_init_args_register
|
||||
def __init__(self, params, learning_rate=0.1, decay=0.9, momentum=0.0, epsilon=1e-10,
|
||||
use_locking=False, centered=False, loss_scale=1.0, weight_decay=0.0):
|
||||
super(RMSProp, self).__init__(learning_rate, params, weight_decay, loss_scale)
|
||||
|
|
|
|||
|
|
@ -19,6 +19,7 @@ from mindspore.common.tensor import Tensor
|
|||
import mindspore.common.dtype as mstype
|
||||
from mindspore._checkparam import Validator as validator
|
||||
from .optimizer import Optimizer
|
||||
from .optimizer import opt_init_args_register
|
||||
|
||||
_sgd_opt = C.MultitypeFuncGraph("sgd_opt")
|
||||
|
||||
|
|
@ -133,6 +134,8 @@ class SGD(Optimizer):
|
|||
>>> loss = nn.SoftmaxCrossEntropyWithLogits()
|
||||
>>> model = Model(net, loss_fn=loss, optimizer=optim)
|
||||
"""
|
||||
|
||||
@opt_init_args_register
|
||||
def __init__(self, params, learning_rate=0.1, momentum=0.0, dampening=0.0, weight_decay=0.0, nesterov=False,
|
||||
loss_scale=1.0):
|
||||
|
||||
|
|
|
|||
|
|
@ -49,13 +49,25 @@ def _init_allreduce_operators(length, split_indices):
|
|||
def _init_allreduce_operators_by_parameters(parameters):
|
||||
""" initialize allreduce communication operators by parameters"""
|
||||
op_list = ()
|
||||
param_fusion = False
|
||||
last_comm_fusion = None
|
||||
first_parameter_flag = True
|
||||
for parameter in parameters:
|
||||
comm_fusion = parameter.comm_fusion
|
||||
if first_parameter_flag:
|
||||
last_comm_fusion = comm_fusion
|
||||
first_parameter_flag = False
|
||||
elif not param_fusion:
|
||||
if comm_fusion != last_comm_fusion:
|
||||
param_fusion = True
|
||||
last_comm_fusion = comm_fusion
|
||||
op = AllReduce('sum', GlobalComm.WORLD_COMM_GROUP)
|
||||
op.add_prim_attr('fusion', comm_fusion)
|
||||
op.add_prim_attr('index', comm_fusion)
|
||||
op_list = op_list + (op,)
|
||||
return op_list
|
||||
if not param_fusion:
|
||||
op_list = ()
|
||||
return op_list, param_fusion
|
||||
|
||||
|
||||
@reduce_opt.register("Tensor", "Bool", "Function", "Function", "Bool", "Tensor")
|
||||
|
|
@ -354,7 +366,7 @@ class DistributedGradReducer(Cell):
|
|||
256.0
|
||||
"""
|
||||
|
||||
def __init__(self, parameters, mean=True, degree=None, fusion_type=1, param_fusion=False):
|
||||
def __init__(self, parameters, mean=True, degree=None, fusion_type=1):
|
||||
super(DistributedGradReducer, self).__init__(auto_prefix=False)
|
||||
self.map_ = C.Map()
|
||||
if degree is None:
|
||||
|
|
@ -371,12 +383,12 @@ class DistributedGradReducer(Cell):
|
|||
if is_parallel_optimizer and split_indices:
|
||||
self.split_fusion = True
|
||||
self.op_list = _init_allreduce_operators(len(parameters), split_indices)
|
||||
elif param_fusion:
|
||||
self.split_fusion = True
|
||||
self.op_list = _init_allreduce_operators_by_parameters(parameters)
|
||||
else:
|
||||
self.split_fusion = False
|
||||
self.allreduce = AllReduce().add_prim_attr('fusion', fusion_type)
|
||||
self.split_fusion = True
|
||||
self.op_list, param_fusion = _init_allreduce_operators_by_parameters(parameters)
|
||||
if not param_fusion:
|
||||
self.split_fusion = False
|
||||
self.allreduce = AllReduce().add_prim_attr('fusion', fusion_type)
|
||||
self.allgather = AllGather(GlobalComm.WORLD_COMM_GROUP)
|
||||
ps_filter = lambda x: x.is_param_ps
|
||||
self.ps_parameters = tuple(ps_filter(x) for x in parameters)
|
||||
|
|
|
|||
|
|
@ -139,6 +139,7 @@ class Model:
|
|||
self._device_number = _get_device_num()
|
||||
self._global_rank = _get_global_rank()
|
||||
self._parameter_broadcast = _get_parameter_broadcast()
|
||||
self._metrics = metrics
|
||||
|
||||
self._check_amp_level_arg(optimizer, amp_level)
|
||||
self._check_for_graph_cell(kwargs)
|
||||
|
|
@ -175,7 +176,7 @@ class Model:
|
|||
|
||||
def _check_kwargs(self, kwargs):
|
||||
for arg in kwargs:
|
||||
if arg not in ['loss_scale_manager', 'keep_batchnorm_fp32', 'total_steps']:
|
||||
if arg not in ['loss_scale_manager', 'keep_batchnorm_fp32']:
|
||||
raise ValueError(f"Unsupported arg '{arg}'")
|
||||
|
||||
def _check_reuse_dataset(self, dataset):
|
||||
|
|
@ -187,15 +188,18 @@ class Model:
|
|||
def _build_acc_network(self, kwargs):
|
||||
"""Build the acc network."""
|
||||
processor = acc.AutoAcc(self._acc_level, kwargs)
|
||||
if processor.level not in ["O1", "O2"]:
|
||||
return
|
||||
if self._optimizer is None:
|
||||
logger.warning("In acc mode, the optimizer must be defined.")
|
||||
return
|
||||
if self._eval_network is None:
|
||||
logger.warning("In acc mode, the eval_network must be defined.")
|
||||
if self._eval_network is None and self._metrics is None:
|
||||
logger.warning("In acc mode, the eval_network and metrics cannot be undefined at the same time.")
|
||||
return
|
||||
|
||||
self._network, self._optimizer = processor.network_auto_process_train(self._network, self._optimizer)
|
||||
self._eval_network = processor.network_auto_process_eval(self._eval_network)
|
||||
if self._eval_network is not None:
|
||||
self._eval_network = processor.network_auto_process_eval(self._eval_network)
|
||||
|
||||
def _build_train_network(self):
|
||||
"""Build train network"""
|
||||
|
|
|
|||
|
|
@ -48,6 +48,7 @@ eval_interval: 1
|
|||
enable_cache: False
|
||||
cache_session_id: ""
|
||||
mode_name: "GRAPH"
|
||||
acc_mode: "O0"
|
||||
|
||||
# Export options
|
||||
device_id: 0
|
||||
|
|
|
|||
|
|
@ -48,6 +48,7 @@ eval_interval: 1
|
|||
enable_cache: False
|
||||
cache_session_id: ""
|
||||
mode_name: "GRAPH"
|
||||
acc_mode: "O0"
|
||||
|
||||
# Export options
|
||||
device_id: 0
|
||||
|
|
|
|||
|
|
@ -50,6 +50,7 @@ eval_interval: 1
|
|||
enable_cache: False
|
||||
cache_session_id: ""
|
||||
mode_name: "GRAPH"
|
||||
acc_mode: "O0"
|
||||
|
||||
# Export options
|
||||
device_id: 0
|
||||
|
|
|
|||
|
|
@ -50,6 +50,7 @@ eval_interval: 1
|
|||
enable_cache: False
|
||||
cache_session_id: ""
|
||||
mode_name: "GRAPH"
|
||||
acc_mode: "O0"
|
||||
|
||||
# Export options
|
||||
device_id: 0
|
||||
|
|
|
|||
|
|
@ -48,6 +48,7 @@ eval_interval: 1
|
|||
enable_cache: False
|
||||
cache_session_id: ""
|
||||
mode_name: "GRAPH"
|
||||
acc_mode: "O0"
|
||||
|
||||
# Export options
|
||||
device_id: 0
|
||||
|
|
|
|||
|
|
@ -0,0 +1,78 @@
|
|||
# Builtin Configurations(DO NOT CHANGE THESE CONFIGURATIONS unless you know exactly what you are doing)
|
||||
enable_modelarts: False
|
||||
# Url for modelarts
|
||||
data_url: ""
|
||||
train_url: ""
|
||||
checkpoint_url: ""
|
||||
# Path for local
|
||||
run_distribute: False
|
||||
enable_profiling: False
|
||||
data_path: "/cache/data"
|
||||
output_path: "/cache/train"
|
||||
load_path: "/cache/checkpoint_path/"
|
||||
device_target: "Ascend"
|
||||
checkpoint_path: "./checkpoint/"
|
||||
checkpoint_file_path: ""
|
||||
|
||||
# ==============================================================================
|
||||
# Training options
|
||||
optimizer: "Momentum"
|
||||
infer_label: ""
|
||||
class_num: 1001
|
||||
batch_size: 256
|
||||
loss_scale: 1024
|
||||
momentum: 0.9
|
||||
weight_decay: 0.0001
|
||||
epoch_size: 90
|
||||
pretrain_epoch_size: 0
|
||||
save_checkpoint: True
|
||||
save_checkpoint_epochs: 5
|
||||
keep_checkpoint_max: 10
|
||||
warmup_epochs: 0
|
||||
lr_decay_mode: "linear"
|
||||
use_label_smooth: True
|
||||
label_smooth_factor: 0.1
|
||||
lr_init: 0
|
||||
lr_max: 0.8
|
||||
lr_end: 0.0
|
||||
|
||||
net_name: "resnet50"
|
||||
dataset: "imagenet2012"
|
||||
device_num: 1
|
||||
pre_trained: ""
|
||||
run_eval: False
|
||||
eval_dataset_path: ""
|
||||
parameter_server: False
|
||||
filter_weight: False
|
||||
save_best_ckpt: True
|
||||
eval_start_epoch: 40
|
||||
eval_interval: 1
|
||||
enable_cache: False
|
||||
cache_session_id: ""
|
||||
mode_name: "GRAPH"
|
||||
acc_mode: "O1"
|
||||
|
||||
# Export options
|
||||
device_id: 0
|
||||
width: 224
|
||||
height: 224
|
||||
file_name: "resnet50"
|
||||
file_format: "AIR"
|
||||
ckpt_file: ""
|
||||
network_dataset: "resnet50_imagenet2012"
|
||||
|
||||
---
|
||||
# Help description for each configuration
|
||||
enable_modelarts: "Whether training on modelarts, default: False"
|
||||
data_url: "Dataset url for obs"
|
||||
checkpoint_url: "The location of checkpoint for obs"
|
||||
data_path: "Dataset path for local"
|
||||
output_path: "Training output path for local"
|
||||
load_path: "The location of checkpoint for obs"
|
||||
device_target: "Target device type, available: [Ascend, GPU, CPU]"
|
||||
enable_profiling: "Whether enable profiling while training, default: False"
|
||||
num_classes: "Class for dataset"
|
||||
batch_size: "Batch size for training and evaluation"
|
||||
epoch_size: "Total training epochs."
|
||||
checkpoint_path: "The location of the checkpoint file."
|
||||
checkpoint_file_path: "The location of the checkpoint file."
|
||||
|
|
@ -51,6 +51,7 @@ eval_interval: 1
|
|||
enable_cache: False
|
||||
cache_session_id: ""
|
||||
mode_name: "GRAPH"
|
||||
acc_mode: "O0"
|
||||
|
||||
# Export options
|
||||
device_id: 0
|
||||
|
|
|
|||
|
|
@ -51,6 +51,7 @@ eval_interval: 1
|
|||
enable_cache: False
|
||||
cache_session_id: ""
|
||||
mode_name: "GRAPH"
|
||||
acc_mode: "O0"
|
||||
|
||||
# Export options
|
||||
device_id: 0
|
||||
|
|
|
|||
|
|
@ -50,6 +50,7 @@ eval_interval: 1
|
|||
enable_cache: False
|
||||
cache_session_id: ""
|
||||
mode_name: "GRAPH"
|
||||
acc_mode: "O0"
|
||||
|
||||
# Export options
|
||||
device_id: 0
|
||||
|
|
|
|||
|
|
@ -25,6 +25,7 @@ eval: False
|
|||
save_ckpt: False
|
||||
mode_name: "GRAPH"
|
||||
dtype: "fp16"
|
||||
acc_mode: "O0"
|
||||
|
||||
# Export options
|
||||
device_id: 0
|
||||
|
|
|
|||
|
|
@ -51,6 +51,7 @@ eval_interval: 1
|
|||
enable_cache: False
|
||||
cache_session_id: ""
|
||||
mode_name: "GRAPH"
|
||||
acc_mode: "O0"
|
||||
|
||||
# Export options
|
||||
device_id: 0
|
||||
|
|
|
|||
|
|
@ -124,9 +124,9 @@ def create_dataset2(dataset_path, do_train, repeat_num=1, batch_size=32, target=
|
|||
device_num = 1
|
||||
|
||||
if device_num == 1:
|
||||
data_set = ds.ImageFolderDataset(dataset_path, num_parallel_workers=8, shuffle=True)
|
||||
data_set = ds.ImageFolderDataset(dataset_path, num_parallel_workers=12, shuffle=True)
|
||||
else:
|
||||
data_set = ds.ImageFolderDataset(dataset_path, num_parallel_workers=8, shuffle=True,
|
||||
data_set = ds.ImageFolderDataset(dataset_path, num_parallel_workers=12, shuffle=True,
|
||||
num_shards=device_num, shard_id=rank_id)
|
||||
|
||||
image_size = 224
|
||||
|
|
@ -152,7 +152,7 @@ def create_dataset2(dataset_path, do_train, repeat_num=1, batch_size=32, target=
|
|||
|
||||
type_cast_op = C2.TypeCast(mstype.int32)
|
||||
|
||||
data_set = data_set.map(operations=trans, input_columns="image", num_parallel_workers=8)
|
||||
data_set = data_set.map(operations=trans, input_columns="image", num_parallel_workers=12)
|
||||
# only enable cache for eval
|
||||
if do_train:
|
||||
enable_cache = False
|
||||
|
|
@ -160,10 +160,10 @@ def create_dataset2(dataset_path, do_train, repeat_num=1, batch_size=32, target=
|
|||
if not cache_session_id:
|
||||
raise ValueError("A cache session_id must be provided to use cache.")
|
||||
eval_cache = ds.DatasetCache(session_id=int(cache_session_id), size=0)
|
||||
data_set = data_set.map(operations=type_cast_op, input_columns="label", num_parallel_workers=8,
|
||||
data_set = data_set.map(operations=type_cast_op, input_columns="label", num_parallel_workers=12,
|
||||
cache=eval_cache)
|
||||
else:
|
||||
data_set = data_set.map(operations=type_cast_op, input_columns="label", num_parallel_workers=8)
|
||||
data_set = data_set.map(operations=type_cast_op, input_columns="label", num_parallel_workers=12)
|
||||
|
||||
# apply batch operations
|
||||
data_set = data_set.batch(batch_size, drop_remainder=True)
|
||||
|
|
|
|||
|
|
@ -105,7 +105,8 @@ def set_parameter():
|
|||
gradients_mean=True)
|
||||
set_algo_parameters(elementwise_op_strategy_follow=True)
|
||||
if config.net_name == "resnet50" or config.net_name == "se-resnet50":
|
||||
context.set_auto_parallel_context(all_reduce_fusion_config=[85, 160])
|
||||
if config.acc_mode not in ["O1", "O2"]:
|
||||
context.set_auto_parallel_context(all_reduce_fusion_config=[85, 160])
|
||||
elif config.net_name == "resnet101":
|
||||
context.set_auto_parallel_context(all_reduce_fusion_config=[80, 210, 313])
|
||||
init()
|
||||
|
|
@ -228,7 +229,8 @@ def train_net():
|
|||
model = Model(net, loss_fn=loss, optimizer=opt, metrics=metrics, eval_network=dist_eval_network)
|
||||
else:
|
||||
model = Model(net, loss_fn=loss, optimizer=opt, loss_scale_manager=loss_scale, metrics=metrics,
|
||||
amp_level="O2", keep_batchnorm_fp32=False, eval_network=dist_eval_network)
|
||||
amp_level="O2", acc_level=config.acc_mode, keep_batchnorm_fp32=False,
|
||||
eval_network=dist_eval_network)
|
||||
|
||||
if config.optimizer == "Thor" and config.dataset == "imagenet2012":
|
||||
from src.lr_generator import get_thor_damping
|
||||
|
|
|
|||
Loading…
Reference in New Issue