MCPcopy Create free account
hub / github.com/pytorch/executorch / _insert_lowered_submodule

Function _insert_lowered_submodule

exir/backend/backend_api.py:198–259  ·  view source on GitHub ↗
(
    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],
)

Source from the content-addressed store, hash-verified

196
197
198def _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,

Calls 9

get_lowered_module_nameFunction · 0.90
generate_debug_handleFunction · 0.85
keysMethod · 0.80
inserting_beforeMethod · 0.80
erase_nodeMethod · 0.80
appendMethod · 0.45
get_attrMethod · 0.45
call_functionMethod · 0.45

Tested by

no test coverage detected