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

Method _ssa_rewrite

caffe2/python/onnx/frontend.py:306–339  ·  view source on GitHub ↗
(cls, net, init_net, value_info)

Source from the content-addressed store, hash-verified

304
305 @classmethod
306 def _ssa_rewrite(cls, net, init_net, value_info):
307 def ssa_name(name, version, version_cnt=None):
308 if version == 0:
309 return name
310 if version_cnt and len(version_cnt.get(name, {})) <= 1:
311 return name
312 return '{}_{}'.format(name, version)
313
314 if init_net:
315 for op in init_net.op:
316 assert re.match('GivenTensor.*Fill', op.type), "type is {}, \n{}".format(op.type, op)
317 assert len(op.output) == 1
318
319 ssa, blob_versions = caffe2_core.get_ssa(net)
320 version_cnt = {}
321 versioned_blobs = []
322 for versioned_input, versioned_output in ssa:
323 versioned_blobs += versioned_input
324 versioned_blobs += versioned_output
325
326 for (name, version) in versioned_blobs:
327 if name not in version_cnt:
328 version_cnt[name] = {version}
329 else:
330 version_cnt[name].add(version)
331
332 assert len(net.op) == len(ssa)
333 for op, (versioned_inputs, versioned_outputs) in zip(net.op, ssa):
334 op.input[:] = [ssa_name(name, version, version_cnt)
335 for name, version in versioned_inputs]
336 op.output[:] = [ssa_name(name, version, version_cnt)
337 for name, version in versioned_outputs]
338 net.external_output[:] = [ssa_name(name, blob_versions[name], version_cnt)
339 for name in net.external_output]
340
341 @classmethod
342 def caffe2_net_to_onnx_model(cls, *args, **kwargs):

Callers 4

ssa_rewriteMethod · 0.80
test_ssaMethod · 0.80
test_idempotenceMethod · 0.80

Calls 4

ssa_nameFunction · 0.50
matchMethod · 0.45
formatMethod · 0.45
addMethod · 0.45

Tested by 2

test_ssaMethod · 0.64
test_idempotenceMethod · 0.64