| 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): |