(gm, subgm_tag, subgm_cb, ptn)
| 523 | |
| 524 | |
| 525 | def _partition_graph_into_submodules(gm, subgm_tag, subgm_cb, ptn): |
| 526 | from torch.fx.passes.utils.fuser_utils import ( |
| 527 | erase_nodes, |
| 528 | fuse_as_graphmodule, |
| 529 | insert_subgm, |
| 530 | legalize_graph, |
| 531 | topo_sort, |
| 532 | ) |
| 533 | |
| 534 | partitions = ptn.propose_partitions() |
| 535 | # insert meta for each partition group |
| 536 | for i, partition in enumerate(partitions): |
| 537 | for node in partition.nodes: |
| 538 | node.meta[subgm_tag] = i |
| 539 | |
| 540 | for i in range(len(partitions)): |
| 541 | # find nodes with same group id in current graph |
| 542 | node_list = [ |
| 543 | node for node in gm.graph.nodes if node.meta.get(subgm_tag, "") == i |
| 544 | ] |
| 545 | # fuse group nodes into submodule |
| 546 | sorted_nodes = topo_sort(node_list) |
| 547 | submodule_name = f"{subgm_tag}_{i}" |
| 548 | subgm, orig_inputs, orig_outputs = fuse_as_graphmodule( |
| 549 | gm, sorted_nodes, submodule_name |
| 550 | ) |
| 551 | # insert submodule & trim group nodes |
| 552 | gm = insert_subgm( |
| 553 | gm, |
| 554 | subgm_cb(subgm, submodule_name), |
| 555 | orig_inputs, |
| 556 | orig_outputs, |
| 557 | ) |
| 558 | erase_nodes(gm, sorted_nodes) |
| 559 | legalize_graph(gm) |
| 560 | |
| 561 | gm.recompile() |
| 562 | return gm |
| 563 | |
| 564 | |
| 565 | def _canonicalize_graph_with_lowered_module(gm, subgm_tag, compiler_specs): |
no test coverage detected