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

Method call_function

exir/pass_base.py:341–397  ·  view source on GitHub ↗
(
            self,
            target: torch.fx.node.Target,
            args: Tuple[Argument, ...],
            kwargs: Dict[str, Argument],
        )

Source from the content-addressed store, hash-verified

339 return self.callback.output(args[0], NodeMetadata(self.node.meta)).data
340
341 def call_function(
342 self,
343 target: torch.fx.node.Target,
344 args: Tuple[Argument, ...],
345 kwargs: Dict[str, Argument],
346 ) -> ProxyValue:
347 meta = NodeMetadata(self.node.meta)
348
349 if target == operator.getitem:
350 value, key = args
351 return self.callback.call_getitem(value, key, meta)
352 elif getattr(target, "__module__", None) in {
353 "_operator",
354 "builtins",
355 "math",
356 }:
357 assert callable(target)
358 return self.callback.call_sym(target, args, meta)
359 elif target in _TORCH_SYM_OPS:
360 assert callable(target)
361 return self.callback.call_sym(target, args, meta)
362 elif isinstance(
363 target, (torch._ops.OpOverload, torch._ops.OpOverloadPacket)
364 ):
365 return self.callback.call_operator(
366 target,
367 args,
368 kwargs,
369 meta,
370 )
371 elif target == torch.ops.higher_order.cond:
372 pred, true_fn, false_fn, inputs = args
373 return self.callback.call_cond(pred, true_fn, false_fn, inputs, meta)
374 elif target == torch.ops.higher_order.while_loop:
375 cond, body, carried_inputs, additional_inputs = args
376 return self.callback.call_while(
377 cond, body, carried_inputs, additional_inputs, meta
378 )
379 elif target == torch.ops.higher_order.map_impl:
380 f, mapped_args, operands = args # type: ignore[assignment]
381 return self.callback.call_map(f, mapped_args, operands, meta)
382 elif target == torch.ops.higher_order.scan:
383 combine_fn, init, xs, additional_inputs = args # type: ignore[assignment]
384 return self.callback.call_scan(
385 combine_fn, init, xs, additional_inputs, meta
386 )
387 # For other unregistered HigherOrderOps, just interpret them blindly
388 elif isinstance(target, torch._ops.HigherOrderOperator):
389 return self.callback._fx(
390 "call_function",
391 target,
392 args,
393 kwargs,
394 meta,
395 )
396 else:
397 raise ExportPassBaseError(f"Unsupported target type: {target}")
398

Callers

nothing calls this directly

Calls 10

NodeMetadataClass · 0.85
ExportPassBaseErrorClass · 0.85
call_symMethod · 0.80
call_whileMethod · 0.80
call_scanMethod · 0.80
call_getitemMethod · 0.45
call_operatorMethod · 0.45
call_condMethod · 0.45
call_mapMethod · 0.45
_fxMethod · 0.45

Tested by

no test coverage detected