(proxy_mode, func_overload, lowered_module, *args)
| 43 | |
| 44 | # pyre-ignore |
| 45 | def trace_call_delegate(proxy_mode, func_overload, lowered_module, *args): |
| 46 | # pyre-ignore |
| 47 | def _unwrap_proxy(e): |
| 48 | if not isinstance(e, (torch.Tensor, torch.SymInt, torch.SymFloat)): |
| 49 | return e |
| 50 | return get_proxy_slot( |
| 51 | cast(torch.Tensor, e), proxy_mode.tracer, e, lambda e: e.proxy |
| 52 | ) |
| 53 | |
| 54 | if not is_lowered_module(lowered_module): |
| 55 | raise ValueError( |
| 56 | "executorch_call_delegate()'s first argument must be a LoweredBackendModule" |
| 57 | ) |
| 58 | |
| 59 | with disable_proxy_modes_tracing(): |
| 60 | out = call_delegate_cpu(lowered_module, *args) |
| 61 | |
| 62 | get_lowered_module_name(proxy_mode.tracer.root, lowered_module) |
| 63 | |
| 64 | node_args = (lowered_module, *args) |
| 65 | proxy_args = pytree.tree_map(_unwrap_proxy, node_args) |
| 66 | out_proxy = proxy_mode.tracer.create_proxy( |
| 67 | "call_function", |
| 68 | func_overload, |
| 69 | proxy_args, |
| 70 | {}, |
| 71 | name="executorch_call_delegate", |
| 72 | ) |
| 73 | return track_tensor_tree( |
| 74 | out, out_proxy, constant=None, tracer=proxy_mode.tracer |
| 75 | ) |
| 76 | |
| 77 | @executorch_call_delegate.py_impl(torch._C.DispatchKey.CompositeExplicitAutograd) |
| 78 | # pyre-ignore |
no test coverage detected