(net, file_name, graph_name="net", op_only=True, blob_rename_func=None)
| 526 | |
| 527 | |
| 528 | def save_graph_base(net, file_name, graph_name="net", op_only=True, blob_rename_func=None): |
| 529 | graph = None |
| 530 | ops = net.op |
| 531 | if blob_rename_func is not None: |
| 532 | ops = _modify_blob_names(ops, blob_rename_func) |
| 533 | if not op_only: |
| 534 | graph = net_drawer.GetPydotGraph(ops, graph_name, rankdir="TB") |
| 535 | else: |
| 536 | graph = net_drawer.GetPydotGraphMinimal( |
| 537 | ops, graph_name, rankdir="TB", minimal_dependency=True |
| 538 | ) |
| 539 | |
| 540 | try: |
| 541 | par_dir = os.path.dirname(file_name) |
| 542 | if not os.path.exists(par_dir): |
| 543 | os.makedirs(par_dir) |
| 544 | |
| 545 | format = os.path.splitext(os.path.basename(file_name))[-1] |
| 546 | if format == ".png": |
| 547 | graph.write_png(file_name) |
| 548 | elif format == ".pdf": |
| 549 | graph.write_pdf(file_name) |
| 550 | elif format == ".svg": |
| 551 | graph.write_svg(file_name) |
| 552 | else: |
| 553 | print("Incorrect format {}".format(format)) |
| 554 | except Exception as e: |
| 555 | print("Error when writing graph to image {}".format(e)) |
| 556 | |
| 557 | return graph |
| 558 | |
| 559 | |
| 560 | # ==== torch/utils_toffee/aten_to_caffe2.py ==================================== |
no test coverage detected