MCPcopy Create free account
hub / github.com/pytorch/pytorch / TranslateModel

Method TranslateModel

caffe2/python/caffe_translator.py:223–279  ·  view source on GitHub ↗
(
        cls,
        caffe_net,
        pretrained_net,
        is_test=False,
        net_state=None,
        remove_legacy_pad=False,
        input_dims=None
    )

Source from the content-addressed store, hash-verified

221
222 @classmethod
223 def TranslateModel(
224 cls,
225 caffe_net,
226 pretrained_net,
227 is_test=False,
228 net_state=None,
229 remove_legacy_pad=False,
230 input_dims=None
231 ):
232 net_state = caffe_pb2.NetState() if net_state is None else net_state
233 net = caffe2_pb2.NetDef()
234 net.name = caffe_net.name
235 net_params = caffe2_pb2.TensorProtos()
236 if len(caffe_net.layers) > 0:
237 raise ValueError(
238 'I think something is wrong. This translation script '
239 'only accepts new style layers that are stored in the '
240 'layer field.'
241 )
242 if not input_dims:
243 input_dims = _GetInputDims(caffe_net)
244 for layer in caffe_net.layer:
245 if not _ShouldInclude(net_state, layer):
246 log.info('Current net state does not need layer {}'
247 .format(layer.name))
248 continue
249 log.info('Translate layer {}'.format(layer.name))
250 # Get pretrained one
251 pretrained_layers = (
252 [l for l in pretrained_net.layer
253 if l.name == layer.name] + [l
254 for l in pretrained_net.layers
255 if l.name == layer.name]
256 )
257 if len(pretrained_layers) > 1:
258 raise ValueError(
259 'huh? more than one pretrained layer of one name?')
260 elif len(pretrained_layers) == 1:
261 pretrained_blobs = [
262 utils.CaffeBlobToNumpyArray(blob)
263 for blob in pretrained_layers[0].blobs
264 ]
265 else:
266 # No pretrained layer for the given layer name. We'll just pass
267 # no parameter blobs.
268 # print 'No pretrained layer for layer', layer.name
269 pretrained_blobs = []
270 operators, params = cls.TranslateLayer(
271 layer, pretrained_blobs, is_test, net=net,
272 net_params=net_params, input_dims=input_dims)
273 net.op.extend(operators)
274 net_params.protos.extend(params)
275 if remove_legacy_pad:
276 assert input_dims, \
277 'Please specify input_dims to remove legacy_pad'
278 net = _RemoveLegacyPad(net, net_params, input_dims)
279 return net, net_params
280

Callers 2

setUpModuleFunction · 0.80
TranslateModelFunction · 0.80

Calls 7

_GetInputDimsFunction · 0.85
_ShouldIncludeFunction · 0.85
_RemoveLegacyPadFunction · 0.85
infoMethod · 0.80
TranslateLayerMethod · 0.80
formatMethod · 0.45
extendMethod · 0.45

Tested by 1

setUpModuleFunction · 0.64