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

Method call_operator

exir/passes/sym_to_tensor_pass.py:27–64  ·  view source on GitHub ↗
(self, op, args, kwargs, meta: NodeMetadata)

Source from the content-addressed store, hash-verified

25
26 # pyre-ignore
27 def call_operator(self, op, args, kwargs, meta: NodeMetadata):
28 # pyre-ignore
29 def is_sym(value, arg) -> bool:
30 if isinstance(value, ProxyValue) and not value.is_tensor():
31 if isinstance(arg.type, torch.TensorType) and type(value.data) in {
32 SymInt,
33 SymFloat,
34 SymBool,
35 }:
36 return True
37 return False
38
39 def corresponding_dtype(
40 symbol: Union[SymInt, SymFloat, SymBool]
41 ) -> torch.dtype:
42 if isinstance(symbol, SymInt):
43 return torch.int32
44 elif isinstance(symbol, SymFloat):
45 return torch.float32
46 elif isinstance(symbol, SymBool):
47 return torch.bool
48 else:
49 raise AssertionError(f"Unsupported data type: {type(symbol)}")
50
51 def try_coerce(value: PyTree, arg: torch.Argument) -> PyTree:
52 if is_sym(value, arg):
53 return self.call_operator(
54 torch.ops.aten.scalar_tensor.default,
55 (value,),
56 {"dtype": corresponding_dtype(value.data)},
57 meta,
58 )
59 else:
60 return value
61
62 args, kwargs = map_args(op, try_coerce, args, kwargs)
63
64 return super().call_operator(op, args, kwargs, meta)

Callers 1

try_coerceMethod · 0.95

Calls 1

map_argsFunction · 0.90

Tested by

no test coverage detected