58 lines
2.2 KiB
Python
58 lines
2.2 KiB
Python
"""
|
|
Copyright (c) 2017-2019 Intel Corporation
|
|
|
|
Licensed under the Apache License, Version 2.0 (the "License");
|
|
you may not use this file except in compliance with the License.
|
|
You may obtain a copy of the License at
|
|
|
|
http://www.apache.org/licenses/LICENSE-2.0
|
|
|
|
Unless required by applicable law or agreed to in writing, software
|
|
distributed under the License is distributed on an "AS IS" BASIS,
|
|
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
See the License for the specific language governing permissions and
|
|
limitations under the License.
|
|
"""
|
|
from mo.front.caffe.collect_attributes import merge_attrs
|
|
from mo.front.caffe.extractors.utils import embed_input
|
|
from mo.front.extractor import FrontExtractorOp
|
|
from mo.ops.op import Op
|
|
|
|
|
|
class DataAugmentationFrontExtractor(FrontExtractorOp):
|
|
op = 'DataAugmentation'
|
|
enabled = True
|
|
|
|
@staticmethod
|
|
def extract(node):
|
|
proto_layer = node.pb
|
|
param = proto_layer.augmentation_param
|
|
# slice_dim is deprecated parameter and is used as alias for axis
|
|
# however if slice_dim is defined and axis is default, we use slice_dim
|
|
update_attrs = {
|
|
'crop_width': param.crop_width,
|
|
'crop_height': param.crop_height,
|
|
'write_augmented': param.write_augmented,
|
|
'max_multiplier': param.max_multiplier,
|
|
'augment_during_test': int(param.augment_during_test),
|
|
'recompute_mean': param.recompute_mean,
|
|
'write_mean': param.write_mean,
|
|
'mean_per_pixel': int(param.mean_per_pixel),
|
|
'mean': param.mean,
|
|
'mode': param.mode,
|
|
'bottomwidth': param.bottomwidth,
|
|
'bottomheight': param.bottomheight,
|
|
'num': param.num,
|
|
'chromatic_eigvec': param.chromatic_eigvec
|
|
}
|
|
|
|
mapping_rule = merge_attrs(param, update_attrs)
|
|
|
|
if node.model_pb:
|
|
for index in range(0, len(node.model_pb.blobs)):
|
|
embed_input(mapping_rule, index + 1, 'custom_{}'.format(index), node.model_pb.blobs[index].data)
|
|
|
|
# update the attributes of the node
|
|
Op.get_op_class_by_name(__class__.op).update_node_stat(node, mapping_rule)
|
|
return __class__.enabled
|