(
submodule_program: ExportedProgram,
owning_program: ExportedProgram,
call_submodule_node: torch.fx.Node,
submodule_output_node: torch.fx.Node,
lowered_module: LoweredBackendModule,
is_submodule: bool,
toplevel_input_specs_to_delete: Dict[str, InputSpec],
toplevel_output_specs_to_delete: Dict[str, OutputSpec],
)
| 196 | |
| 197 | |
| 198 | def _insert_lowered_submodule( |
| 199 | submodule_program: ExportedProgram, |
| 200 | owning_program: ExportedProgram, |
| 201 | call_submodule_node: torch.fx.Node, |
| 202 | submodule_output_node: torch.fx.Node, |
| 203 | lowered_module: LoweredBackendModule, |
| 204 | is_submodule: bool, |
| 205 | toplevel_input_specs_to_delete: Dict[str, InputSpec], |
| 206 | toplevel_output_specs_to_delete: Dict[str, OutputSpec], |
| 207 | ): |
| 208 | owning_graph_module = call_submodule_node.graph.owning_module |
| 209 | # call delegate args should only use user_inputs |
| 210 | call_delegate_args = [] |
| 211 | # names of input_specs to delete |
| 212 | input_specs_to_delete = toplevel_input_specs_to_delete |
| 213 | # Delete owned constants from the call_submodule_node args |
| 214 | for call_sm_input in call_submodule_node.args: |
| 215 | if ( |
| 216 | isinstance(call_sm_input, torch.fx.Node) |
| 217 | and call_sm_input.name in input_specs_to_delete.keys() |
| 218 | ): |
| 219 | continue |
| 220 | call_delegate_args.append(call_sm_input) |
| 221 | |
| 222 | def generate_debug_handle(ep: ExportedProgram) -> int: |
| 223 | """ |
| 224 | Generate a debug handle for the given ExportedProgram. |
| 225 | """ |
| 226 | debug_handle = 0 |
| 227 | for node in ep.graph_module.graph.nodes: |
| 228 | debug_handle = max(debug_handle, node.meta.get("debug_handle", 0)) |
| 229 | return debug_handle + 1 |
| 230 | |
| 231 | # Replace the partitioned submodule with a lowered submodule |
| 232 | # Add call_method node with function "forward" |
| 233 | with owning_graph_module.graph.inserting_before(call_submodule_node): |
| 234 | lowered_name = get_lowered_module_name(owning_graph_module, lowered_module) |
| 235 | lowered_node = owning_graph_module.graph.get_attr(lowered_name) |
| 236 | call_delegate_node = owning_graph_module.graph.call_function( |
| 237 | executorch_call_delegate, |
| 238 | (lowered_node,) + tuple(call_delegate_args), |
| 239 | call_submodule_node.kwargs, |
| 240 | ) |
| 241 | call_delegate_node.meta["debug_handle"] = generate_debug_handle(owning_program) |
| 242 | call_delegate_node.meta["val"] = [ |
| 243 | out_arg.meta["val"] for out_arg in submodule_output_node.args[0] |
| 244 | ] |
| 245 | call_submodule_node.replace_all_uses_with(call_delegate_node) |
| 246 | owning_graph_module.graph.erase_node(call_submodule_node) |
| 247 | if is_submodule: |
| 248 | assert len(toplevel_input_specs_to_delete) == 0 |
| 249 | assert len(toplevel_output_specs_to_delete) == 0 |
| 250 | elif ( |
| 251 | len(toplevel_input_specs_to_delete) > 0 |
| 252 | or len(toplevel_output_specs_to_delete) > 0 |
| 253 | ): |
| 254 | _unsafe_adjust_original_program( |
| 255 | owning_program, |
no test coverage detected