4
from textwrap import dedent
5
from typing import Any, List, NamedTuple, Optional, Tuple
7
from torch._C import ErrorReport
8
from torch._C._jit_tree_views import SourceRangeFactory
11
def get_source_lines_and_file(
13
error_msg: Optional[str] = None,
14
) -> Tuple[List[str], int, Optional[str]]:
16
Wrapper around inspect.getsourcelines and inspect.getsourcefile.
18
Returns: (sourcelines, file_lino, filename)
22
filename = inspect.getsourcefile(obj)
23
sourcelines, file_lineno = inspect.getsourcelines(obj)
26
f"Can't get source for {obj}. TorchScript requires source access in "
27
"order to carry out compilation, make sure original .py files are "
31
msg += "\n" + error_msg
32
raise OSError(msg) from e
34
return sourcelines, file_lineno, filename
37
def normalize_source_lines(sourcelines: List[str]) -> List[str]:
39
This helper function accepts a list of source lines. It finds the
40
indentation level of the function definition (`def`), then it indents
41
all lines in the function body to a point at or greater than that
42
level. This allows for comments and continued string literals that
43
are at a lower indentation than the rest of the code.
45
sourcelines: function source code, separated into lines by
48
A list of source lines that have been correctly aligned
51
def remove_prefix(text, prefix):
52
return text[text.startswith(prefix) and len(prefix) :]
56
for i, l in enumerate(sourcelines):
57
if l.lstrip().startswith("def"):
68
fn_def = sourcelines[idx]
69
whitespace = fn_def.split("def")[0]
73
whitespace + remove_prefix(s, whitespace) for s in sourcelines[:idx]
76
whitespace + remove_prefix(s, whitespace) for s in sourcelines[idx + 1 :]
80
aligned_prefix.append(fn_def)
81
return aligned_prefix + aligned_suffix
86
class SourceContext(SourceRangeFactory):
92
leading_whitespace_len,
93
uses_true_division=True,
96
super().__init__(source, filename, file_lineno, leading_whitespace_len)
97
self.uses_true_division = uses_true_division
98
self.filename = filename
99
self.funcname = funcname
102
@functools.lru_cache(maxsize=None)
103
def make_source_context(*args):
104
return SourceContext(*args)
108
return SourceContext("", None, 0, 0).make_raw_range(0, 1)
111
class ParsedDef(NamedTuple):
115
filename: Optional[str]
120
sourcelines, file_lineno, filename = get_source_lines_and_file(
121
fn, ErrorReport.call_stack()
123
sourcelines = normalize_source_lines(sourcelines)
124
source = "".join(sourcelines)
125
dedent_src = dedent(source)
126
py_ast = ast.parse(dedent_src)
127
if len(py_ast.body) != 1 or not isinstance(py_ast.body[0], ast.FunctionDef):
129
f"Expected a single top-level function: {filename}:{file_lineno}"
131
leading_whitespace_len = len(source.split("\n", 1)[0]) - len(
132
dedent_src.split("\n", 1)[0]
134
ctx = make_source_context(
135
source, filename, file_lineno, leading_whitespace_len, True, fn.__name__
137
return ParsedDef(py_ast, ctx, source, filename, file_lineno)