(
self,
target: torch.fx.node.Target,
args: Tuple[Argument, ...],
kwargs: Dict[str, Argument],
)
| 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 |
nothing calls this directly
no test coverage detected