This is an automated email from the ASF dual-hosted git repository. tqchen pushed a commit to branch tvmscript-generic-parser-builder in repository https://gitbox.apache.org/repos/asf/tvm.git
commit fbc75bed933136e823ee5824819aa325f62df070 Author: Tianqi Chen <[email protected]> AuthorDate: Mon Sep 21 00:33:20 2026 +0000 Translate original Python syntax through registered construction protocols --- python/tvm/relax/script/builder_v2/__init__.py | 100 ++++- python/tvm/relax/script/v2.py | 50 +++ python/tvm/script/ir_builder/protocol.py | 13 + python/tvm/script/parser_v2/__init__.py | 34 ++ python/tvm/script/parser_v2/annotations.py | 356 +++++++++++++++ python/tvm/script/parser_v2/frontend.py | 486 ++++++++++++++++++++ python/tvm/script/parser_v2/functions.py | 224 ++++++++++ python/tvm/script/parser_v2/ir.py | 28 ++ python/tvm/script/parser_v2/transform.py | 589 +++++++++++++++++++++++++ python/tvm/tirx/script/builder_v2/__init__.py | 79 +++- python/tvm/tirx/script/v2.py | 44 ++ 11 files changed, 2001 insertions(+), 2 deletions(-) diff --git a/python/tvm/relax/script/builder_v2/__init__.py b/python/tvm/relax/script/builder_v2/__init__.py index 7ea3a135da..3b1f493cca 100644 --- a/python/tvm/relax/script/builder_v2/__init__.py +++ b/python/tvm/relax/script/builder_v2/__init__.py @@ -24,6 +24,11 @@ import tvm_ffi as _ffi from tvm import ir as _ir from tvm import relax as _relax +from tvm import tirx as _tir +from tvm.relax.distributed import DeviceMesh as _DeviceMesh +from tvm.relax.distributed import DTensorType as _DTensorType +from tvm.relax.distributed import Placement as _Placement +from tvm.relax.distributed import device_mesh as device_mesh from tvm.script.ir_builder import IRBuilder as _IRBuilder from tvm.script.ir_builder import ir as _I from tvm.script.ir_builder import protocol as _protocol @@ -31,7 +36,9 @@ from tvm.script.ir_builder import protocol as _protocol from .. import builder as _legacy from ..builder import * from ..builder import _ffi_api +from ..builder import distributed as dist from ..builder import frame as _frame +from ..builder.distributed.ir import _lookup_device_mesh @_protocol.expression_args("shape", introduce=True, dtype="int64", scalar_strings=False) @@ -45,6 +52,21 @@ def Tensor(shape=None, dtype=None, vdevice=None, ndim=-1, *, span=None): return _relax.TensorType(shape, dtype, vdevice, ndim, span) +@_protocol.expression_args("shape", introduce=True, dtype="int64", scalar_strings=False) +def DTensor(shape=None, dtype=None, device_mesh=None, placement="", *, ndim=-1, span=None): + """Construct a concrete distributed tensor type with resolved shape symbols.""" + if device_mesh is None: + device_mesh = _DeviceMesh([], _ir.Range(0, 1)) + elif isinstance(device_mesh, _python.str): + device_mesh = _lookup_device_mesh(device_mesh) + if isinstance(placement, _python.str): + placement = _Placement.from_text(placement) + return _DTensorType(Tensor(shape, dtype, ndim=ndim), device_mesh, placement, span) + + +Range = _ir.Range + + @_protocol.expression_args("values", introduce=True, dtype="int64") def Shape(values=None, ndim=-1, *, span=None): """Construct a concrete shape type.""" @@ -93,6 +115,9 @@ def Prim(dtype, *, span=None): return _ir.PrimType(dtype) +Prim.__tvm_parameter_dtype__ = "dtype" + + def Object(*, span=None): """Construct the unconstrained Relax value type.""" return _relax.AnyType(span) @@ -101,6 +126,9 @@ def Object(*, span=None): Any = Object +is_type_var = _ir.is_prim_var + + def type_var(name, *, dtype=None, span=None): """Construct a signature symbol under Relax's default shape dtype policy.""" return _ir.Var(name, "int64" if dtype is None else dtype, span) @@ -222,9 +250,25 @@ def bind_( name_span=None, previous=_protocol.MISSING, declaration=False, + frame_value=False, ): - """Emit an immutable Relax binding and return the newly bound value.""" + """Emit a binding, or name and preserve an existing frame-owned value.""" _check_unterminated() + if frame_value: + if isinstance(value, _python.list | _python.tuple | _ir.Array): + for index, item in enumerate(value): + bind_( + item, + name=None if name is None else f"{name}_{index}", + span=span, + name_span=name_span, + frame_value=True, + ) + elif isinstance(value, _ir.Var): + if name is not None: + _IRBuilder.name(name, value) + _protocol.at(name_span if name_span is not None else span, value) + return value if declaration: if not _ir.is_prim_var(value): raise TypeError("A symbol declaration requires a concrete primitive variable") @@ -321,9 +365,11 @@ __all__ = [ *_legacy.ir.__all__, "Any", "Callable", + "DTensor", "For", "Object", "Prim", + "Range", "Shape", "Tensor", "Tuple", @@ -332,10 +378,62 @@ __all__ = [ "continue_", "bind_", "decl_function", + "device_mesh", + "dist", "emit_", + "is_type_var", "match_cast", "return_", "setitem", "type_var", "unpack", ] + + +def _logical_pair(lhs, rhs, operation, primitive, python_operation): + if not isinstance(lhs, _ir.Expr) and not isinstance(rhs, _ir.Expr): + return python_operation(lhs, rhs) + if _ir.is_prim_expr(lhs) or _ir.is_prim_expr(rhs): + return primitive(lhs, rhs) + return operation(_value(lhs), _value(rhs)) + + +def logical_and(*values): + """Construct conjunction of concrete tensor or primitive expressions.""" + if not values: + raise TypeError("logical_and requires at least one operand") + result = values[0] + for value in values[1:]: + result = _logical_pair(result, value, _relax.op.logical_and, _tir.And, lambda a, b: a and b) + return result + + +def logical_or(*values): + """Construct disjunction of concrete tensor or primitive expressions.""" + if not values: + raise TypeError("logical_or requires at least one operand") + result = values[0] + for value in values[1:]: + result = _logical_pair(result, value, _relax.op.logical_or, _tir.Or, lambda a, b: a or b) + return result + + +def logical_not(value): + """Construct negation without coercing an IR expression to Python bool.""" + if _ir.is_prim_expr(value): + return _tir.Not(value) + if isinstance(value, _ir.Expr): + return _relax.op.logical_not(value) + return not value + + +def select(condition, true_value, false_value): + """Construct an elementwise conditional or select ordinary Python values.""" + if _ir.is_prim_expr(condition): + return _tir.Select(condition, true_value, false_value) + if isinstance(condition, _ir.Expr): + return _relax.op.where(condition, _value(true_value), _value(false_value)) + return true_value if condition else false_value + + +__all__ += ["logical_and", "logical_not", "logical_or", "select"] diff --git a/python/tvm/relax/script/v2.py b/python/tvm/relax/script/v2.py new file mode 100644 index 0000000000..9030ca0d2d --- /dev/null +++ b/python/tvm/relax/script/v2.py @@ -0,0 +1,50 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +"""Opt-in TVMScript entry point using concrete Relax construction operations.""" + +# pylint: disable=wildcard-import,unused-wildcard-import,redefined-builtin +import sys as _sys + +from tvm import relax as _relax +from tvm.script.parser_v2.frontend import make_decorator as _make_decorator +from tvm.script.parser_v2.frontend import make_helper as _make_helper +from tvm.script.parser_v2.frontend import register_namespace as _register_namespace +from tvm.script.parser_v2.functions import register_opaque_factory as _register_opaque_factory + +from . import builder_v2 as _builder +from .builder_v2 import * # noqa: F403 + +function = _make_decorator(_builder, option_map={"pure": "is_pure", "private": "is_private"}) +macro = _make_helper(_builder, preserve_return=True) + +_register_namespace("R", _sys.modules[__name__]) +_register_namespace("relax", _sys.modules[__name__]) + + +def _opaque_function(name, function, source, span): + return _relax.ExternFunc(name, span=span).with_attrs( + { + "is_pyfunc": True, + "function_type": "python", + "python_function_name": name, + "python_source": source, + "python_packed_func": function, + } + ) + + +_register_opaque_factory(_opaque_function) diff --git a/python/tvm/script/ir_builder/protocol.py b/python/tvm/script/ir_builder/protocol.py index 87cb155862..fb4a73cb5a 100644 --- a/python/tvm/script/ir_builder/protocol.py +++ b/python/tvm/script/ir_builder/protocol.py @@ -21,6 +21,7 @@ and return concrete values. Dialects register their own function kinds here, so translation consumes construction policies without importing their owners. """ +from builtins import locals as locals from builtins import slice as slice from contextlib import nullcontext from dataclasses import dataclass @@ -118,3 +119,15 @@ def at(span, value): _at = at + + +def frame_result(frame): + """Read finalized lexical exports without imposing a dialect policy.""" + return getattr(frame, "result", {}) + + +def require_defined(value, name): + """Reject reads of names absent from a finalized lexical export set.""" + if value is MISSING: + raise NameError(f"name {name!r} is not defined") + return value diff --git a/python/tvm/script/parser_v2/__init__.py b/python/tvm/script/parser_v2/__init__.py new file mode 100644 index 0000000000..3128291335 --- /dev/null +++ b/python/tvm/script/parser_v2/__init__.py @@ -0,0 +1,34 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +"""Translate original Python ASTs into calls on a registered construction namespace. + +Shared I infrastructure owns builder lifetime, spans and module identities; each +context selects X from reverse-registered function metadata. Calls return concrete +values, constructor metadata describes syntax, and compiled AST locations remain +those of the original source. Entry modules register policies; this parser never +imports their namespaces. +""" + +from .frontend import ( + from_source, + ir_module, + make_decorator, + make_helper, + parse, + pyfunc, + register_namespace, +) diff --git a/python/tvm/script/parser_v2/annotations.py b/python/tvm/script/parser_v2/annotations.py new file mode 100644 index 0000000000..06b5045cc3 --- /dev/null +++ b/python/tvm/script/parser_v2/annotations.py @@ -0,0 +1,356 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +"""Annotation evaluation and callable-owned expression-argument syntax. + +Constructors receive concrete values. This scope prepares designated symbols, +rewrites only registered expression fields, and evaluates each annotation once. +""" + +import ast +import builtins +import copy +import inspect +import linecache +import re +from typing import TypeVar + + +class AnnotationScope: + """A function's canonical symbols and once-evaluated annotation values.""" + + def __init__(self, env, builder, filename, span): + self.env = dict(env) + self.builder = builder + self.filename = filename + self.span = span + self.symbols = {} + self._evaluated = {} + self._prepared = {} + self._type_vars = {} + self._counter = 0 + self._used_names = set(env) + + def _error(self, node, message): + raise SyntaxError( + message, + ( + self.filename, + node.lineno, + node.col_offset + 1, + linecache.getline(self.filename, node.lineno), + ), + ) + + def _symbol(self, name, node, dtype=None, *, shadow=False): + if not shadow and name in self.symbols: + return self.symbols[name] + value = self.env.get(name) + if not shadow and name in self.env and not isinstance(value, TypeVar): + if getattr(self.builder, "is_type_var", lambda value: False)(value): + self.symbols[name] = value + return value + value = self.builder.type_var(name, dtype=dtype, span=self.span(node)) + self.symbols[name] = self.env[name] = value + return value + + def _canonical_type_var(self, name, node, dtype=None): + value = self.env.get(name) + if not isinstance(value, TypeVar): + return + if value.__constraints__ or value.__bound__ is not None: + self._error(node, "A symbolic TypeVar cannot have constraints or a bound") + if value.__name__ != name: + self._error(node, "A symbolic TypeVar binding must match its declared name") + if value not in self._type_vars: + self._type_vars[value] = self._symbol(name, node, dtype) + else: + self.symbols[name] = self.env[name] = self._type_vars[value] + + def prepare_type_params(self, type_params): + """Introduce explicit host-supported PEP 695 parameters, shadowing outer names.""" + type_var_node = getattr(ast, "TypeVar", ()) + for parameter in type_params: + if not isinstance(parameter, type_var_node): + self._error(parameter, "Only scalar TypeVar parameters are supported") + bound = getattr(parameter, "bound", None) + if bound is not None and not ( + isinstance(bound, ast.Name) + and self.env.get(bound.id, getattr(builtins, bound.id, None)) is int + ): + self._error(parameter, "A symbolic TypeVar bound must be int") + if getattr(parameter, "default_value", None) is not None: + self._error(parameter, "A symbolic TypeVar cannot have a default") + self._symbol(parameter.name, parameter, shadow=True) + + def _resolve(self, node): + # Inspect callable identity without executing the annotation or its arguments. + if isinstance(node, ast.Name): + return self.env.get(node.id, getattr(builtins, node.id, None)) + if isinstance(node, ast.Attribute): + owner = self._resolve(node.value) + if owner is None: + return None + value = inspect.getattr_static(owner, node.attr, None) + if isinstance(value, staticmethod): + return value.__func__ + # Descriptor execution belongs to evaluation, not metadata discovery. + return None if isinstance(value, property) else value + return None + + @staticmethod + def _arguments(call, constructor): + try: + parameters = list(inspect.signature(constructor).parameters.values()) + except (TypeError, ValueError): + return {} + positional = [ + p + for p in parameters + if p.kind + in (inspect.Parameter.POSITIONAL_ONLY, inspect.Parameter.POSITIONAL_OR_KEYWORD) + ] + fields = { + parameter.name: value + for parameter, value in zip(positional, call.args) + if not isinstance(value, ast.Starred) + } + fields.update( + {keyword.arg: keyword.value for keyword in call.keywords if keyword.arg is not None} + ) + return fields + + def _cache_expression(self, node): + name = f"__tvm_annotation_value_{self._counter}" + self._counter += 1 + while name in self.env or name in self._used_names: + name = f"__tvm_annotation_value_{self._counter}" + self._counter += 1 + self.env[name] = self._eval(node) + return ast.copy_location(ast.Name(name, ast.Load()), node) + + def prepare_parameters(self, arguments): + """Preallocate scalar parameter identities before dependent annotations. + + Scalar constructors own dtype metadata. Cached dtype expressions are + substituted in the annotation, so even a dynamic dtype is evaluated once. + ``symbols`` contains scalar parameter objects for the signature builder. + The caller evaluates annotations and registers parameters sequentially, + making each actual parameter available to subsequent annotations. + """ + parameters = [*arguments.posonlyargs, *arguments.args, *arguments.kwonlyargs] + self._used_names.update( + node.id for node in ast.walk(arguments) if isinstance(node, ast.Name) + ) + self._used_names.update(parameter.arg for parameter in parameters) + for parameter in parameters: + annotation = parameter.annotation + if annotation is None: + continue + prepared = copy.deepcopy(annotation) + if isinstance(prepared, ast.Constant) and isinstance(prepared.value, str): + prepared = self._string_expression(prepared) + self._prepared[id(annotation)] = prepared + constructor = self._resolve( + prepared.func if isinstance(prepared, ast.Call) else prepared + ) + declaration = getattr(constructor, "__tvm_declaration_args__", None) + dtype_field = getattr(constructor, "__tvm_parameter_dtype__", None) + if declaration is not None: + dtype = declaration.dtype + elif dtype_field is not None and isinstance(prepared, ast.Call): + dtype_node = self._arguments(prepared, constructor).get(dtype_field) + if dtype_node is None: + dtype = inspect.signature(constructor).parameters[dtype_field].default + if dtype is inspect.Parameter.empty: + continue + else: + cached = self._cache_expression(dtype_node) + dtype = self.env[cached.id] + for index, argument in enumerate(prepared.args): + if argument is dtype_node: + prepared.args[index] = cached + for keyword in prepared.keywords: + if keyword.value is dtype_node: + keyword.value = cached + else: + continue + self._symbol(parameter.arg, parameter, dtype, shadow=True) + return self.symbols + + def _eval(self, node): + expression = ast.Expression(body=node) + ast.fix_missing_locations(expression) + return eval(compile(expression, self.filename, "eval"), self.env) # pylint: disable=eval-used + + def evaluate(self, node, introduce=True): + """Evaluate an annotation once in the prepared signature scope.""" + key = id(node) + if key not in self._evaluated: + prepared = self._prepared.get(key, node) + # A quoted whole annotation is ordinary Python annotation syntax. + if isinstance(prepared, ast.Constant) and isinstance(prepared.value, str): + prepared = self._string_expression(prepared) + self._evaluated[key] = self._eval(self.rewrite(prepared, introduce=introduce)) + return self._evaluated[key] + + def _string_expression(self, node): + try: + expression = ast.parse(node.value, mode="eval").body + except SyntaxError as error: + self._error(node, f"Invalid annotation expression: {error.msg}") + source = "".join(linecache.getlines(self.filename)) + literal = ast.get_source_segment(source, node) if source else None + positions = self._literal_positions(literal, node) if literal else None + lines = node.value.splitlines(keepends=True) + for inner in ast.walk(expression): + if not hasattr(inner, "lineno"): + continue + for line_field, column_field in ( + ("lineno", "col_offset"), + ("end_lineno", "end_col_offset"), + ): + line, column = getattr(inner, line_field), getattr(inner, column_field) + offset = len("".join(lines[: line - 1]).encode("utf-8")) + column + if positions is not None and offset in positions: + line, column = positions[offset] + else: + column += node.col_offset + 1 if line == 1 else 0 + line += node.lineno - 1 + setattr(inner, line_field, line) + setattr(inner, column_field, column) + return expression + + @staticmethod + def _literal_positions(literal, node): + """Map decoded expression byte offsets back through the literal's escapes.""" + match = re.match("(?i:([rub]*))([\"'])", literal) + if match is None: + return None + prefix, quote = match.groups() + width = 3 if literal[len(prefix) :].startswith(quote * 3) else 1 + start, stop = len(prefix) + width, len(literal) - width + delimiter = quote * width + positions, decoded, offset = {}, "", 0 + index = start + + def location(raw_index): + before = literal[:raw_index] + line = node.lineno + before.count("\n") + column = len(before.rsplit("\n", 1)[-1].encode("utf-8")) + return line, column + (node.col_offset if line == node.lineno else 0) + + while index < stop: + end = index + 1 + if literal[index] == "\\" and "r" not in prefix.lower(): + escape = re.match( + r"\\(?:N\{[^}]*\}|u[0-9a-fA-F]{4}|U[0-9a-fA-F]{8}|x[0-9a-fA-F]{2}|[0-7]{1,3}|\r?\n|.)", + literal[index:stop], + ) + if escape: + end = index + len(escape.group()) + piece = literal[index:end] + try: + value = ast.literal_eval(prefix + delimiter + piece + delimiter) + except (SyntaxError, ValueError): + return None + positions[offset] = location(index) + for char in value: + offset += len(char.encode("utf-8")) + positions[offset] = location(end) + decoded += value + index = end + return positions if decoded == node.value else None + + def rewrite(self, node, *, introduce=False): + """Return a copied expression AST, registering new symbols in ``env``. + + Construction code must execute with this scope's updated environment. + No assignment to a generated local is needed for newly introduced names. + """ + scope = self + + class Rewrite(ast.NodeTransformer): + def __init__(self): + self.allow_names = False + self.dtype = None + + def visit_Name(self, current): + if isinstance(current.ctx, ast.Load): + if isinstance(scope.env.get(current.id), TypeVar) and not introduce: + scope._error( + current, "A TypeVar must be introduced in a signature or match scope" + ) + scope._canonical_type_var(current.id, current, self.dtype) + if self.allow_names and ( + current.id in scope.env or not hasattr(builtins, current.id) + ): + scope._symbol(current.id, current, self.dtype) + return current + + def visit_Attribute(self, current): + old_allow = self.allow_names + self.allow_names = False + current.value = self.visit(current.value) + self.allow_names = old_allow + return current + + def expression_field(self, current, metadata, *, nested=False): + old_allow, old_dtype = self.allow_names, self.dtype + self.allow_names = introduce and metadata.introduce + self.dtype = metadata.dtype + try: + if isinstance(current, ast.List | ast.Tuple): + current.elts = [ + self.expression_field(value, metadata, nested=True) + for value in current.elts + ] + return current + if isinstance(current, ast.Constant) and isinstance(current.value, str): + if nested or metadata.scalar_strings: + current = scope._string_expression(current) + return self.visit(current) + finally: + self.allow_names, self.dtype = old_allow, old_dtype + + def visit_Call(self, current): + constructor = scope._resolve(current.func) + metadata = getattr(constructor, "__tvm_expression_args__", None) + fields = scope._arguments(current, constructor) if metadata else {} + marked = ( + {id(value) for name, value in fields.items() if name in metadata.fields} + if metadata + else set() + ) + old_allow = self.allow_names + self.allow_names = False + current.func = self.visit(current.func) + self.allow_names = old_allow + current.args = [ + self.expression_field(value, metadata) + if id(value) in marked + else self.visit(value) + for value in current.args + ] + for keyword in current.keywords: + keyword.value = ( + self.expression_field(keyword.value, metadata) + if id(keyword.value) in marked + else self.visit(keyword.value) + ) + return current + + return ast.fix_missing_locations(Rewrite().visit(copy.deepcopy(node))) diff --git a/python/tvm/script/parser_v2/frontend.py b/python/tvm/script/parser_v2/frontend.py new file mode 100644 index 0000000000..f96fa413f2 --- /dev/null +++ b/python/tvm/script/parser_v2/frontend.py @@ -0,0 +1,486 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +"""Source acquisition and declaration/body execution for registered builders.""" + +import ast +import copy +import inspect +import linecache +import textwrap +from dataclasses import dataclass, field +from functools import wraps +from types import SimpleNamespace + +from tvm import ir +from tvm.script.ir_builder import IRBuilder, protocol +from tvm.script.ir_builder import ir as I + +from .annotations import AnnotationScope +from .functions import FunctionGroup, attach_python, declare_python, is_python_function + +_NAMESPACES = {} + + +def register_namespace(alias, namespace): + """Let an entry module supply a source-level namespace without reverse imports.""" + _NAMESPACES[alias] = namespace + + +def _resolve(node, env, filename): + expression = ast.Expression(copy.deepcopy(node)) + return eval(compile(ast.fix_missing_locations(expression), filename, "eval"), env) + + +def _capture(obj): + target = obj if inspect.isfunction(obj) else None + env = dict(getattr(target, "__globals__", {})) + if target is not None: + closure = inspect.getclosurevars(target) + env.update(closure.globals) + env.update(closure.nonlocals) + filename = inspect.getsourcefile(obj) + # Deferred annotations may be the only use of an enclosing local, so Python + # need not put that value in the function's closure cells. + frames = inspect.stack() + try: + for info in reversed(frames): + if info.filename == filename: + env.update(info.frame.f_locals) + finally: + del frames + if inspect.isclass(obj): + env.update(vars(obj)) + return env + + +def _inside_class(function): + frame = inspect.currentframe().f_back + try: + while frame is not None: + local = frame.f_locals + if local.get("__module__") == function.__module__ and "__qualname__" in local: + return True + if frame.f_code.co_filename == function.__code__.co_filename: + return False + frame = frame.f_back + finally: + del frame + return False + + +def make_decorator(builder, *, option_map=None, defaults=None): + """Create a parsing decorator whose construction policy belongs to its caller.""" + mapping, default_options = dict(option_map or {}), dict(defaults or {}) + + def decorator(function=None, **options): + def apply(function): + function.__tvm_function_kind__ = decorator.__tvm_function_kind__ + function.__tvm_function_options__ = options + if _inside_class(function): + return function + return parse(function, _capture(function)) + + return apply(function) if function is not None else apply + + return protocol.register_function( + decorator, builder, option_map=mapping, defaults=default_options + ) + + +def make_helper(builder, *, preserve_return=True): + """Create an explicit construction helper using the active shared builder.""" + + def decorator(function=None, **options): + def apply(function): + @wraps(function) + def invoke(*args, **kwargs): + bound = inspect.signature(function).bind(*args, **kwargs) + bound.apply_defaults() + compiler = Compiler(function, _capture(function)) + node = compiler.tree.body[0] + return compiler.run_statements( + node.body, + builder, + {**compiler.env, **bound.arguments}, + set(bound.arguments), + preserve_return=preserve_return, + ) + + invoke.__tvm_construction_helper__ = (builder, options) + return invoke + + return apply(function) if function is not None else apply + + return decorator + + +def pyfunc(function): + """Mark a function whose body and execution remain ordinary Python.""" + function.__tvm_python_function__ = True + return function + + +protocol.register_function(pyfunc, None, python=True) + + +@dataclass +class Signature: + node: ast.FunctionDef + kind: protocol.FunctionKind + options: dict + scope: AnnotationScope + params: dict = field(default_factory=dict) + result_type: object = protocol.MISSING + reference: object = None + + +class Compiler: + """Build from a copied original AST, retaining file and range information.""" + + def __init__(self, source, env=None, filename=None): + self.env = {**_NAMESPACES, **(env or {})} + self.original = source + if isinstance(source, str): + text = source + self.filename = filename or "<tvmscript>" + start, indent = 1, 0 + linecache.cache[self.filename] = ( + len(text), + None, + text.splitlines(keepends=True), + self.filename, + ) + else: + lines, start = inspect.getsourcelines(source) + text = "".join(lines) + self.filename = filename or inspect.getsourcefile(source) + indent = len(lines[0]) - len(lines[0].lstrip()) + self.tree = ast.parse(textwrap.dedent(text), self.filename) + if start != 1: + ast.increment_lineno(self.tree, start - 1) + if indent: + for node in ast.walk(self.tree): + if hasattr(node, "col_offset"): + node.col_offset += indent + node.end_col_offset += indent + self.used_names = {n.id for n in ast.walk(self.tree) if isinstance(n, ast.Name)} + self.used_names.update(n.arg for n in ast.walk(self.tree) if isinstance(n, ast.arg)) + self.used_names.update( + n.name for n in ast.walk(self.tree) if isinstance(n, ast.FunctionDef | ast.ClassDef) + ) + self.used_names.update(self.env) + self.counter = 0 + self.function_kinds = {} + self.spans = [] + self.span_indices = {} + self.source_name = ir.SourceName(self.filename) + self.span_name = self.fresh("spans") + self.builder_name = self.fresh("builder") + self.infrastructure_name = self.fresh("infrastructure") + + def fresh(self, prefix): + while True: + self.counter += 1 + name = f"__script_{prefix}_{self.counter}" + if name not in self.used_names: + self.used_names.add(name) + return name + + def span(self, node): + key = (node.lineno, node.end_lineno, node.col_offset, node.end_col_offset) + if key not in self.span_indices: + self.span_indices[key] = len(self.spans) + self.spans.append(ir.Span(self.source_name, *key)) + return self.spans[self.span_indices[key]] + + def span_ast(self, node): + self.span(node) + index = self.span_indices[ + (node.lineno, node.end_lineno, node.col_offset, node.end_col_offset) + ] + return ast.copy_location( + ast.Subscript(ast.Name(self.span_name, ast.Load()), ast.Constant(index), ast.Load()), + node, + ) + + def function_kind(self, node, env): + if id(node) in self.function_kinds: + return self.function_kinds[id(node)] + for decorator in node.decorator_list: + target = decorator.func if isinstance(decorator, ast.Call) else decorator + value = _resolve(target, env, self.filename) + kind = protocol.function_kind(value) + if kind is not None: + options = dict(kind.metadata.get("defaults", {})) + if isinstance(decorator, ast.Call): + if decorator.args: + raise SyntaxError("Function decorators accept keyword options only") + for item in decorator.keywords: + if item.arg is None: + options.update(_resolve(item.value, env, self.filename)) + else: + options[item.arg] = _resolve(item.value, env, self.filename) + mapping = kind.metadata.get("option_map", {}) + options = { + mapping.get(key, key): value + for key, value in options.items() + if key != "check_well_formed" + } + self.function_kinds[id(node)] = (kind, options) + return kind, options + raise SyntaxError(f"Function {node.name!r} has no registered construction kind") + + def declare(self, node, env, *, local=False): + kind, options = self.function_kind(node, env) + scope = AnnotationScope(env, kind.builder, self.filename, self.span) + spec = Signature(node, kind, options, scope) + if kind.metadata.get("python"): + return spec + X = kind.builder + mode = {"local": True} if local else {} + with X.decl_function(**options, **mode, span=self.span(node)) as frame: + X.func_name(node.name) + scope.prepare_type_params(getattr(node, "type_params", [])) + scope.prepare_parameters(node.args) + if node.args.posonlyargs or node.args.kwonlyargs or node.args.vararg or node.args.kwarg: + raise SyntaxError("IR signatures require ordinary named parameters") + for argument in node.args.args: + if argument.annotation is None: + raise SyntaxError(f"Parameter {argument.arg!r} requires an annotation") + annotation = scope.evaluate(argument.annotation) + value = X.arg( + argument.arg, + scope.symbols.get(argument.arg, annotation), + span=self.span(argument), + ) + spec.params[argument.arg] = scope.env[argument.arg] = value + if node.returns is not None: + spec.result_type = scope.evaluate(node.returns) + X.func_ret_type(spec.result_type) + spec.reference = frame.reference + scope.env[node.name] = spec.reference + return spec + + def define(self, spec, env, *, local=False): + X = spec.kind.builder + mode = {"local": True, "reference": spec.reference} if local else {} + with X.function(**spec.options, **mode, span=self.span(spec.node)) as frame: + X.func_name(spec.node.name) + for name, value in spec.params.items(): + X.arg(name, value) + if spec.result_type is not protocol.MISSING: + X.func_ret_type(spec.result_type) + scope_env = {**env, **spec.scope.env, **spec.params} + self.run_statements( + spec.node.body, + X, + scope_env, + set(spec.params) | set(spec.scope.symbols), + scope=spec.scope, + ) + return frame.function + + def run_statements(self, body, builder, env, bound_names, *, scope=None, preserve_return=False): + from .transform import Transformer + + namespace = dict(env) + namespace.update( + { + self.builder_name: builder, + self.infrastructure_name: protocol, + self.span_name: self.spans, + } + ) + nested = {} + nested_name = self.fresh("nested") + + def nested_statement(node): + if is_python_function(self, node, namespace): + return [copy.deepcopy(node)] + index = len(nested) + nested[index] = node + value = ast.Call( + ast.Name(nested_name, ast.Load()), + [ + ast.Constant(index), + ast.Call( + ast.Attribute( + ast.Name(self.infrastructure_name, ast.Load()), "locals", ast.Load() + ), + [], + [], + ), + ], + [], + ) + return [ast.copy_location(ast.Assign([ast.Name(node.name, ast.Store())], value), node)] + + def build_nested(index, values): + node = nested[index] + local_env = {**namespace, **values} + group = FunctionGroup(self, [node], local_env, local=True) + return group.define(node.name) + + namespace[nested_name] = build_nested + injected = {} + + def rewrite_expression(node): + node = scope.rewrite(node) if scope is not None else copy.deepcopy(node) + if isinstance(node, ast.Call): + target = scope._resolve(node.func) if scope is not None else None + try: + replacement = getattr(builder, "__tvm_call_overrides__", {}).get(target) + except TypeError: + replacement = None + if replacement is not None: + name = self.fresh("call") + injected[name] = replacement + node.func = ast.copy_location(ast.Name(name, ast.Load()), node.func) + method = None + values = [] + if isinstance(node, ast.BoolOp): + method = "logical_and" if isinstance(node.op, ast.And) else "logical_or" + values = node.values + elif isinstance(node, ast.UnaryOp) and isinstance(node.op, ast.Not): + method, values = "logical_not", [node.operand] + elif isinstance(node, ast.IfExp): + method, values = "select", [node.test, node.body, node.orelse] + if method is not None: + node = ast.copy_location( + ast.Call( + ast.Attribute(ast.Name(self.builder_name, ast.Load()), method, ast.Load()), + values, + [], + ), + node, + ) + return node + + signature_values = {} + if scope is not None: + for name, value in scope.symbols.items(): + alias = self.fresh("symbol") + namespace[alias] = value + signature_values[name] = ast.Name(alias, ast.Load()) + + transformer = Transformer( + filename=self.filename, + environment=namespace, + builder_name=self.builder_name, + infrastructure_name=self.infrastructure_name, + span=self.span_ast, + fresh=self.fresh, + signature_names=set(bound_names), + signature_values=signature_values, + expression_rewriter=rewrite_expression, + nested_function=nested_statement, + preserve_return=preserve_return, + ) + statements = transformer.transform_statements(copy.deepcopy(body)) + namespace.update(injected) + if scope is not None: + namespace.update(scope.env) + bound_names = set(bound_names) | set(scope.symbols) + names = sorted(name for name in bound_names if name in namespace) + helper_name = self.fresh("body") + helper = ast.FunctionDef( + name=helper_name, + args=ast.arguments( + posonlyargs=[], + args=[ast.arg(name) for name in names], + kwonlyargs=[], + kw_defaults=[], + defaults=[], + ), + body=statements or [ast.Pass()], + decorator_list=[], + ) + ast.copy_location(helper, body[0]) + module = ast.fix_missing_locations(ast.Module([helper], [])) + exec(compile(module, self.filename, "exec"), namespace) + return namespace[helper_name](*(namespace[name] for name in names)) + + def build(self): + nodes = self.tree.body + if len(nodes) == 1 and isinstance(nodes[0], ast.ClassDef): + root = nodes[0] + statements = root.body + elif len(nodes) == 1 and isinstance(nodes[0], ast.FunctionDef): + root = None + statements = nodes + else: + raise SyntaxError("Source must contain one function or module class") + env = dict(self.env) + functions = [node for node in statements if isinstance(node, ast.FunctionDef)] + python_functions = [] + with IRBuilder() as builder: + with I.ir_module(): + references = {node.name: I.reserve_function(node.name) for node in functions} + env.update(references) + if root is not None: + env[root.name] = SimpleNamespace(**references) + for statement in statements: + if not isinstance(statement, ast.FunctionDef): + exec( + compile( + ast.fix_missing_locations( + ast.Module([copy.deepcopy(statement)], []) + ), + self.filename, + "exec", + ), + env, + ) + ir_functions = [] + for node in functions: + if is_python_function(self, node, env): + original = self.env.get(node.name) + original = ( + original + if getattr(original, "__tvm_python_function__", False) + else None + ) + record = declare_python(self, node, env, original=original) + python_functions.append(record) + references[node.name] = env[node.name] = record.reference + else: + ir_functions.append(node) + group = FunctionGroup(self, ir_functions, env) + results = group.define_all() + module = builder.get() + if python_functions: + attach_python(module, python_functions) + return module if root is not None else results[functions[0].name] + + +def parse(source, extra_vars=None, *, filename=None, **options): + """Construct from source using only entry-module registered construction policies.""" + env = {} if isinstance(source, str) else _capture(source) + env.update(extra_vars or {}) + return Compiler(source, env, filename).build() + + +def ir_module(module=None, **options): + """Build a module after its registered function signatures have been declared.""" + + def apply(module): + return parse(module, _capture(module), **options) + + return apply(module) if module is not None else apply + + +from_source = parse diff --git a/python/tvm/script/parser_v2/functions.py b/python/tvm/script/parser_v2/functions.py new file mode 100644 index 0000000000..540fbea6b1 --- /dev/null +++ b/python/tvm/script/parser_v2/functions.py @@ -0,0 +1,224 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +"""Function-group construction and opaque Python member registration.""" + +import ast +import copy +import dis +import inspect +import linecache +from dataclasses import dataclass +from types import CodeType + +from tvm.script.ir_builder import ir as I + +_OPAQUE_FACTORY = None + + +def register_opaque_factory(factory): + """Register the owner-provided constructor for a module's opaque function slot. + + The factory receives ``(name, python_callable, source_text, span)`` and must + return a concrete BaseFunc. Registration never executes the Python body. + """ + global _OPAQUE_FACTORY + if not callable(factory): + raise TypeError("An opaque function factory must be callable") + _OPAQUE_FACTORY = factory + + +def is_python_function(compiler, node, env): + """Identify a Python function through its registered function-kind metadata.""" + kind, _ = compiler.function_kind(node, env) + return bool(kind.metadata.get("python")) + + +def _error(compiler, node, message): + raise SyntaxError( + message, + ( + compiler.filename, + node.lineno, + node.col_offset + 1, + linecache.getline(compiler.filename, node.lineno), + ), + ) + + +def materialize_python(compiler, node, env, *, original=None): + """Retain an existing Python callable, or execute its unchanged definition. + + Source-only module members share ``env``, preserving ordinary global lookup. + Nested Python definitions should remain directly in the construction helper's + AST, so Python itself creates their lexical closure cells. + """ + if original is None and not isinstance(compiler.original, str): + candidate = env.get(node.name) + if getattr(candidate, "__tvm_python_function__", False): + original = candidate + if original is not None: + if not callable(original): + _error(compiler, node, "A Python function definition must retain a callable") + return original + module = ast.fix_missing_locations(ast.Module([copy.deepcopy(node)], [])) + exec( + compile( + module, + compiler.filename, + "exec", + flags=getattr(compiler, "compile_flags", 0), + dont_inherit=True, + ), + env, + ) + return env[node.name] + + +@dataclass(frozen=True) +class PythonFunction: + """Original Python callable and its separate opaque module reference.""" + + name: str + function: object + reference: object + + +def declare_python(compiler, node, env, *, original=None): + """Construct a real opaque module entry and retain its Python implementation.""" + if _OPAQUE_FACTORY is None: + _error(compiler, node, "No opaque Python function constructor has been registered") + function = materialize_python(compiler, node, env, original=original) + try: + source = inspect.getsource(function) + except (OSError, TypeError): + source = ast.unparse(node) + opaque = _OPAQUE_FACTORY(node.name, function, source, compiler.span(node)) + reference = I.decl_function(node.name, opaque) + I.def_function(node.name, opaque) + return PythonFunction(node.name, function, reference) + + +def attach_python(module, functions): + """Attach Python callables to a completed module with opaque entries. + + This preserves construction and Python execution. Device conversion, method + binding, and runtime registration remain responsibilities of the runtime's + module adapter; attaching callables alone does not establish those bridges. + """ + existing = dict(getattr(module, "pyfuncs", {})) + for function in functions: + existing[function.name] = function.function + module.pyfuncs = existing + return module + + +def _global_loads(node, filename): + """Find free sibling uses without confusing attributes or shadowed locals.""" + node = copy.deepcopy(node) + node.decorator_list = [] + node.returns = None + for argument in (*node.args.posonlyargs, *node.args.args, *node.args.kwonlyargs): + argument.annotation = None + module = ast.fix_missing_locations(ast.Module([node], [])) + code = compile(module, filename, "exec", dont_inherit=True) + names = set() + + def visit(current): + for instruction in dis.get_instructions(current): + if instruction.opname in ("LOAD_GLOBAL", "LOAD_NAME"): + names.add(instruction.argval) + for value in current.co_consts: + if isinstance(value, CodeType): + visit(value) + + # Skip the module's decoration/default expressions: only body references + # determine whether a local definition uses a still-undefined peer. + for value in code.co_consts: + if isinstance(value, CodeType) and value.co_name == node.name: + visit(value) + return names + + +class FunctionGroup: + """Declare sibling IR signatures before defining any of their bodies. + + An optional ``reserve(name)`` callback allocates stable identities before + annotation evaluation, when the enclosing module or builder supports it. + Local definitions retain source ordering and reject forward sibling uses; + self-recursion uses the function's own completed declaration. + """ + + def __init__(self, compiler, nodes, env, *, local=False, reserve=None): + self.compiler = compiler + self.local = local + self.env = dict(env) + self.signatures = {} + self.references = {} + self.results = {} + nodes = list(nodes) + names = {node.name for node in nodes} + if len(names) != len(nodes): + duplicate = next( + node + for index, node in enumerate(nodes) + if node.name in {previous.name for previous in nodes[:index]} + ) + _error(compiler, duplicate, f"Duplicate function declaration {duplicate.name!r}") + if reserve is not None: + self.references.update((node.name, reserve(node.name)) for node in nodes) + self.env.update(self.references) + for node in nodes: + signature = compiler.declare(node, self.env, local=local) + self.signatures[node.name] = signature + self.references[node.name] = signature.reference + self.env[node.name] = signature.reference + for signature in self.signatures.values(): + shadowed = signature.params.keys() | signature.scope.symbols.keys() + signature.scope.env.update( + (name, reference) + for name, reference in self.references.items() + if name not in shadowed + ) + + def define(self, name, env=None): + """Define one declared body, returning the cached construction reference.""" + signature = self.signatures[name] + if name in self.results: + _error(self.compiler, signature.node, f"Function {name!r} is already defined") + if self.local: + unresolved = ( + (_global_loads(signature.node, self.compiler.filename) & self.references.keys()) + - self.results.keys() + - {name} + ) + if unresolved: + _error( + self.compiler, + signature.node, + f"Local function {name!r} refers to undefined sibling " + f"{sorted(unresolved)[0]!r}; " + "mutually recursive local definitions are unsupported", + ) + visible = {**self.env, **(env or {}), **self.references} + self.results[name] = self.compiler.define(signature, visible, local=self.local) + return self.references[name] + + def define_all(self): + """Define bodies in source order after every signature has been registered.""" + for name in self.signatures: + self.define(name) + return self.results diff --git a/python/tvm/script/parser_v2/ir.py b/python/tvm/script/parser_v2/ir.py new file mode 100644 index 0000000000..9471ada2a0 --- /dev/null +++ b/python/tvm/script/parser_v2/ir.py @@ -0,0 +1,28 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +"""Shared module entry paired with explicit construction exports.""" + +import sys + +from tvm.ir import Range as Range +from tvm.script.ir_builder.ir import * # noqa: F403 + +from .frontend import ir_module as ir_module +from .frontend import pyfunc as pyfunc +from .frontend import register_namespace + +register_namespace("I", sys.modules[__name__]) diff --git a/python/tvm/script/parser_v2/transform.py b/python/tvm/script/parser_v2/transform.py new file mode 100644 index 0000000000..ab8729c0df --- /dev/null +++ b/python/tvm/script/parser_v2/transform.py @@ -0,0 +1,589 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +"""Translate Python syntax into calls on a context-selected construction namespace. + +The translator owns ordering, lexical scopes, and original source locations. +Construction operations and callable syntax metadata come from its environment; +ordinary expressions remain nested Python expressions producing concrete values. +""" + +import ast +import copy +import inspect + + +class Transformer(ast.NodeTransformer): + """Lower statements without importing or discovering construction namespaces.""" + + def __init__( + self, + filename, + environment, + builder_name, + infrastructure_name, + span, + fresh, + signature_names=None, + expression_rewriter=None, + nested_function=None, + preserve_return=False, + signature_values=None, + ): + self.filename = filename + self.environment = dict(environment) + self.builder_name = builder_name + self.infrastructure_name = infrastructure_name + self.span = span + self.fresh = fresh + self.signature_names = set(signature_names or ()) + self.signature_values = dict(signature_values or {}) + self.bound = set(self.signature_names) + self.optional = {} + self.expression_rewriter = expression_rewriter + self.nested_function = nested_function + self.preserve_return = preserve_return + + def transform_statements(self, body): + """Transform a copy, retaining the caller's original source tree.""" + result = [] + for statement in copy.deepcopy(body): + translated = self.visit(statement) + if translated is not None: + block = translated if isinstance(translated, list) else [translated] + context = self._call( + self.infrastructure_name, "span_context", [self.span(statement)], statement + ) + result.append(self._with(context, block, statement)) + return result + + def _error(self, node, message): + raise SyntaxError(message, (self.filename, node.lineno, node.col_offset + 1, None)) + + @staticmethod + def _located(value, original): + return ast.copy_location(value, original) + + def _name(self, name, original, store=False): + return self._located(ast.Name(name, ast.Store() if store else ast.Load()), original) + + def _attribute(self, namespace, member, original): + return self._located( + ast.Attribute(self._name(namespace, original), member, ast.Load()), original + ) + + def _call(self, namespace, member, args, original, **keywords): + return self._located( + ast.Call( + self._attribute(namespace, member, original), + args, + [ast.keyword(arg=key, value=value) for key, value in keywords.items()], + ), + original, + ) + + def _operation(self, member, args, original, **keywords): + return self._call( + self.builder_name, member, args, original, span=self.span(original), **keywords + ) + + def _statement(self, expression, original): + return self._located(ast.Expr(expression), original) + + def _assign(self, name, value, original): + return self._located(ast.Assign([self._name(name, original, True)], value), original) + + def _cache(self, value, original, prefix="value"): + name = self.fresh(prefix) + return self._assign(name, value, original), self._name(name, original) + + def _resolve(self, node): + if isinstance(node, ast.Name): + return self.environment.get(node.id) + if isinstance(node, ast.Attribute): + owner = self._resolve(node.value) + if owner is None: + return None + value = inspect.getattr_static(owner, node.attr, None) + if isinstance(value, staticmethod): + return value.__func__ + if inspect.isfunction(value): + return ( + value + if inspect.ismodule(owner) or inspect.isclass(owner) + else value.__get__(owner) + ) + # Inspect syntax metadata without executing a source-level descriptor. + if hasattr(type(value), "__get__"): + return None + return value + return None + + def _is_declaration(self, node): + if not isinstance(node, ast.Call): + return False + constructor = self._resolve(node.func) + policy = getattr(constructor, "__tvm_declaration_args__", None) + if policy is None or any(isinstance(arg, ast.Starred) for arg in node.args): + return False + if any(keyword.arg is None for keyword in node.keywords): + return False + try: + arguments = inspect.signature(constructor).bind_partial( + *node.args, **{keyword.arg: keyword.value for keyword in node.keywords} + ) + except (TypeError, ValueError): + return False + value = arguments.arguments.get(policy.value_parameter) + return value is None or (isinstance(value, ast.Constant) and value.value is None) + + def _expression(self, original, *, attach_span=True): + node = copy.deepcopy(original) + if isinstance(getattr(node, "ctx", None), ast.Store): + return node + if isinstance(node, ast.Name) and node.id in self.optional: + value = self._call( + self.infrastructure_name, + "require_defined", + [copy.deepcopy(self.optional[node.id]), ast.Constant(node.id)], + node, + ) + return self._call( + self.infrastructure_name, "_at", [self.span(original), value], original + ) + if self.expression_rewriter is not None: + node = self.expression_rewriter(node) + if isinstance(node, ast.Await | ast.Yield | ast.YieldFrom | ast.NamedExpr): + self._error(original, f"Unsupported expression: {type(node).__name__}") + if isinstance(node, ast.JoinedStr): + # JoinedStr's children must remain literal fragments/FormattedValue nodes. + for child in node.values: + if isinstance(child, ast.FormattedValue): + child.value = self._expression(child.value) + if child.format_spec is not None: + child.format_spec = self._format_spec(child.format_spec) + else: + self._expression_children(node) + if isinstance(node, ast.Starred) or ( + isinstance(node, ast.Name) and isinstance(node.ctx, ast.Store) + ): + return node + if not attach_span: + return node + return self._call(self.infrastructure_name, "_at", [self.span(original), node], original) + + def _format_spec(self, node): + for child in node.values: + if isinstance(child, ast.FormattedValue): + child.value = self._expression(child.value) + if child.format_spec is not None: + child.format_spec = self._format_spec(child.format_spec) + return node + + def _expression_children(self, node): + for field, value in ast.iter_fields(node): + if isinstance(value, ast.expr): + setattr(node, field, self._expression(value)) + elif isinstance(value, list): + for index, item in enumerate(value): + if isinstance(item, ast.expr): + value[index] = self._expression(item) + elif isinstance(item, ast.AST): + self._expression_children(item) + elif isinstance(value, ast.AST): + self._expression_children(value) + + def _index(self, node): + if isinstance(node, ast.Slice): + fields = [ + self._expression(value) if value is not None else ast.Constant(None) + for value in (node.lower, node.upper, node.step) + ] + return self._call(self.infrastructure_name, "slice", fields, node) + if isinstance(node, ast.Tuple): + return self._located( + ast.Tuple([self._index(value) for value in node.elts], ast.Load()), node + ) + return self._expression(node) + + def _bind(self, target, value, statement, ty=None, declaration=False, frame_value=False): + if isinstance(target, ast.Name): + keywords = {"name": ast.Constant(target.id), "name_span": self.span(target)} + if frame_value: + keywords["frame_value"] = ast.Constant(True) + if ty is not None: + keywords["ty"] = ty + if declaration: + keywords["declaration"] = ast.Constant(True) + if target.id in self.signature_names: + keywords["previous"] = copy.deepcopy( + self.signature_values.get(target.id, self._name(target.id, target)) + ) + elif target.id in self.bound: + keywords["previous"] = self._name(target.id, target) + elif target.id in self.optional: + keywords["previous"] = copy.deepcopy(self.optional[target.id]) + self.bound.add(target.id) + self.optional.pop(target.id, None) + return [ + self._assign( + target.id, self._operation("bind_", [value], statement, **keywords), target + ) + ] + if isinstance(target, ast.Subscript): + return [ + self._statement( + self._operation( + "setitem", + [self._expression(target.value), self._index(target.slice), value], + statement, + ), + statement, + ) + ] + if isinstance(target, ast.Tuple | ast.List): + # Each Python unpack finishes before visiting that level's targets. A nested + # unpack occurs only when reached, preserving assignment and failure order. + names = [self.fresh("unpack") for _ in target.elts] + pattern = [] + for item, name in zip(target.elts, names): + temporary = self._name(name, item, True) + pattern.append( + self._located(ast.Starred(temporary, ast.Store()), item) + if isinstance(item, ast.Starred) + else temporary + ) + unpack = self._call(self.builder_name, "unpack", [value], target) + assignment = self._located( + ast.Assign([ast.Tuple(pattern, ast.Store())], unpack), target + ) + result = [assignment] + for item, name in zip(target.elts, names): + item = item.value if isinstance(item, ast.Starred) else item + result.extend( + self._bind(item, self._name(name, item), statement, frame_value=frame_value) + ) + return result + self._error(target, f"Unsupported assignment target: {type(target).__name__}") + + def visit_Assign(self, node): + # Cache first: stores evaluate RHS before target base/index, and chained + # assignments share exactly one RHS evaluation. + declaration = self._is_declaration(node.value) + cache, value = self._cache( + self._expression(node.value, attach_span=not declaration), node.value + ) + result = [cache] + resolved = self._resolve(node.value) + for target in node.targets: + result.extend(self._bind(target, copy.deepcopy(value), node, declaration=declaration)) + if isinstance(target, ast.Name): + if resolved is not None: + self.environment[target.id] = resolved + else: + self.environment.pop(target.id, None) + return result + + def visit_AnnAssign(self, node): + if not isinstance(node.target, ast.Name): + self._error(node.target, "An annotated binding requires a name") + result = [] + if node.value is None: + value = self._attribute(self.infrastructure_name, "MISSING", node) + else: + cache, value = self._cache( + self._expression(node.value, attach_span=not self._is_declaration(node.value)), + node.value, + ) + result.append(cache) + result.extend( + self._bind( + node.target, + value, + node, + self._expression(node.annotation), + self._is_declaration(node.value), + ) + ) + return result + + def visit_AugAssign(self, node): + result = [] + if isinstance(node.target, ast.Name): + old, previous = self._cache( + self._expression(self._located(ast.Name(node.target.id, ast.Load()), node.target)), + node.target, + "old", + ) + result.append(old) + value = self._located(ast.BinOp(previous, node.op, self._expression(node.value)), node) + value = self._call(self.infrastructure_name, "_at", [self.span(node), value], node) + return result + self._bind(node.target, value, node) + if not isinstance(node.target, ast.Subscript): + self._error(node.target, "An augmented assignment requires a name or index") + base_stmt, base = self._cache( + self._expression(node.target.value), node.target.value, "base" + ) + key_stmt, key = self._cache(self._index(node.target.slice), node.target.slice, "key") + load = self._located( + ast.Subscript(copy.deepcopy(base), copy.deepcopy(key), ast.Load()), node.target + ) + old_stmt, old = self._cache( + self._call( + self.infrastructure_name, "_at", [self.span(node.target), load], node.target + ), + node.target, + "old", + ) + result.extend([base_stmt, key_stmt, old_stmt]) + value = self._located(ast.BinOp(old, node.op, self._expression(node.value)), node) + value = self._call(self.infrastructure_name, "_at", [self.span(node), value], node) + result.append(self._statement(self._operation("setitem", [base, key, value], node), node)) + return result + + def visit_Expr(self, node): + return self._statement(self._operation("emit_", [self._expression(node.value)], node), node) + + def visit_Return(self, node): + value = self._expression(node.value) if node.value is not None else ast.Constant(None) + if self.preserve_return: + return self._located(ast.Return(value), node) + return self._statement(self._operation("return_", [value], node), node) + + def visit_Break(self, node): + return self._statement(self._operation("break_", [], node), node) + + def visit_Continue(self, node): + return self._statement(self._operation("continue_", [], node), node) + + def visit_Assert(self, node): + message = self._expression(node.msg) if node.msg is not None else ast.Constant("") + return self._statement( + self._operation("assert_", [self._expression(node.test), message], node), node + ) + + @staticmethod + def _assigned_names(body): + names = set() + + class Names(ast.NodeVisitor): + def visit_Name(self, node): + if isinstance(node.ctx, ast.Store): + names.add(node.id) + + def visit_FunctionDef(self, node): + names.add(node.name) + + visit_AsyncFunctionDef = visit_FunctionDef + + def visit_Lambda(self, node): + pass + + visitor = Names() + for statement in body: + visitor.visit(statement) + return names + + def _scope(self, body, original, prefix, initial=None): + outer_bound, outer_environment, outer_optional = self.bound, self.environment, self.optional + referenced = { + node.id + for statement in body + for node in ast.walk(statement) + if isinstance(node, ast.Name) + } + captures = sorted((outer_bound | outer_optional.keys()).intersection(referenced)) + self.bound, self.environment = set(outer_bound), dict(outer_environment) + self.optional = { + name: self._name(name, original) + if name in captures and not self.preserve_return + else value + for name, value in outer_optional.items() + } + prefix_statements = [] if initial is None else initial() + translated = prefix_statements + self.transform_statements(body) + self.bound, self.environment, self.optional = outer_bound, outer_environment, outer_optional + if self.preserve_return: + return translated or [self._located(ast.Pass(), original)] + # A helper isolates construction locals. Optional exports are captured as + # values or MISSING, and checked only when the original body reads them. + defaults = [ + self._name(name, original) + if name in outer_bound + else copy.deepcopy(outer_optional[name]) + for name in captures + ] + helper = self.fresh(prefix) + arguments = ast.arguments( + posonlyargs=[], + args=[ast.arg(arg=name) for name in captures], + vararg=None, + kwonlyargs=[], + kw_defaults=[], + kwarg=None, + defaults=defaults, + ) + definition = self._located( + ast.FunctionDef( + helper, arguments, translated or [self._located(ast.Pass(), original)], [], None + ), + original, + ) + if "type_params" in ast.FunctionDef._fields: + definition.type_params = [] + invocation = self._located(ast.Call(self._name(helper, original), [], []), original) + return [definition, self._statement(invocation, original)] + + def _exports(self, frame, candidates, original): + mapping_stmt, mapping = self._cache( + self._call( + self.infrastructure_name, "frame_result", [self._name(frame, original)], original + ), + original, + "exports", + ) + result = [mapping_stmt] + for name in sorted(candidates): + key = ast.Constant(name) + condition = self._located( + ast.Compare(copy.deepcopy(key), [ast.In()], [copy.deepcopy(mapping)]), original + ) + value = self._located(ast.Subscript(copy.deepcopy(mapping), key, ast.Load()), original) + result.append( + self._located( + ast.If(condition, [self._assign(name, value, original)], []), original + ) + ) + for name in candidates - self.bound: + self.optional[name] = self._located( + ast.Call( + ast.Attribute(copy.deepcopy(mapping), "get", ast.Load()), + [ + ast.Constant(name), + self._attribute(self.infrastructure_name, "MISSING", original), + ], + [], + ), + original, + ) + return result + + def _with(self, context, body, original, target=None): + return self._located( + ast.With([ast.withitem(context, target)], body or [ast.Pass()]), original + ) + + def visit_If(self, node): + frame = self.fresh("conditional") + branches = [ + self._with( + self._operation("Then", [], node), self._scope(node.body, node, "then"), node + ) + ] + if node.orelse: + branches.append( + self._with( + self._operation("Else", [], node), self._scope(node.orelse, node, "else"), node + ) + ) + region = self._with( + self._operation("If", [self._expression(node.test)], node), + branches, + node, + self._name(frame, node, True), + ) + return [region, *self._exports(frame, self._assigned_names(node.body + node.orelse), node)] + + def visit_For(self, node): + if node.orelse: + self._error(node, "A construction loop does not support an else clause") + frame, values = self.fresh("loop"), self.fresh("indices") + context = self._assign( + frame, self._operation("For", [self._expression(node.iter)], node), node + ) + body = self._scope( + node.body, + node, + "body", + lambda: self._bind_entered(node.target, self._name(values, node.target), node), + ) + region = self._with(self._name(frame, node), body, node, self._name(values, node, True)) + return [context, region, *self._exports(frame, self._assigned_names(node.body), node)] + + def visit_While(self, node): + if node.orelse: + self._error(node, "A construction loop does not support an else clause") + frame = self.fresh("loop") + region = self._with( + self._operation("While", [self._expression(node.test)], node), + self._scope(node.body, node, "body"), + node, + self._name(frame, node, True), + ) + return [region, *self._exports(frame, self._assigned_names(node.body), node)] + + def _bind_entered(self, target, value, original): + # Context/iteration targets introduce lexical names; they never reassign + # an outer mutable variable merely because its source spelling matches. + for item in ast.walk(target): + if isinstance(item, ast.Name): + self.bound.discard(item.id) + self.optional.pop(item.id, None) + return self._bind(target, value, original, frame_value=True) + + def visit_With(self, node): + item = node.items[0] + body = node.body + if len(node.items) > 1: + nested = self._located(ast.With(node.items[1:], node.body), node) + body = [nested] + manager, value = self.fresh("context"), self.fresh("entered") + cache = self._assign(manager, self._expression(item.context_expr), item.context_expr) + initial = ( + None + if item.optional_vars is None + else lambda: self._bind_entered( + item.optional_vars, self._name(value, item.context_expr), node + ) + ) + region = self._with( + self._name(manager, node), + self._scope(body, node, "scope", initial), + node, + self._name(value, node, True), + ) + return [cache, region, *self._exports(manager, self._assigned_names(body), node)] + + def visit_FunctionDef(self, node): + self.bound.add(node.name) + if self.nested_function is not None: + return self.nested_function(node) + # Ordinary construction helpers retain normal Python execution. Registered + # function-kind entry points are resolved by the enclosing compiler callback. + if any( + getattr(self._resolve(decorator), "__tvm_function_kind__", None) + for decorator in node.decorator_list + ): + self._error(node, "A registered nested function requires a function compiler") + return node + + def visit_Pass(self, node): + return node + + def generic_visit(self, node): + if isinstance(node, ast.stmt): + self._error(node, f"Unsupported statement: {type(node).__name__}") + return super().generic_visit(node) diff --git a/python/tvm/tirx/script/builder_v2/__init__.py b/python/tvm/tirx/script/builder_v2/__init__.py index 1647786a15..de1fa0b758 100644 --- a/python/tvm/tirx/script/builder_v2/__init__.py +++ b/python/tvm/tirx/script/builder_v2/__init__.py @@ -37,6 +37,8 @@ from tvm.tirx.script.builder import * # pylint: disable=wildcard-import,unused- from tvm.tirx.script.builder import _ffi_api from tvm.tirx.script.builder import frame as _frame +is_type_var = _ir.is_prim_var + def type_var(name, *, dtype=None, span=None): """Construct a signature symbol; shape symbols default to int64.""" @@ -217,11 +219,27 @@ def bind_( name_span=None, previous=_MISSING, declaration=False, + frame_value=False, ): - """Bind concrete values, preserving existing mutable scalar storage.""" + """Bind values, or name a frame-owned value without introducing new storage.""" name_span = span if name_span is None else name_span _check_unterminated() with _span_context(span): + if frame_value: + if isinstance(value, _python.list | _python.tuple | _ir.Array): + for index, item in enumerate(value): + bind_( + item, + name=None if name is None else f"{name}_{index}", + span=span, + name_span=name_span, + frame_value=True, + ) + elif isinstance(value, _ir.Var | _tir.IterVar | _tir.Layout): + _name(value, name, name_span) + elif isinstance(value, _ir.TensorLoad) and _tir.is_buffer_var(value.source): + _name(value.source, name, name_span) + return value if declaration: if not _ir.is_prim_var(value): raise TypeError("A symbol declaration requires a concrete primitive variable") @@ -445,3 +463,62 @@ for _constructor in vars(_T).values(): if isinstance(_constructor, _T.DtypeConstructor): _register_declaration(_constructor, dtype=_constructor._dtype_str) del _constructor + + +def range_(*args): + """Construct a serial loop from the source builtin range arguments.""" + if len(args) == 1: + start, stop, step = 0, args[0], 1 + elif len(args) == 2: + start, stop = args + step = 1 + elif len(args) == 3: + start, stop, step = args + else: + raise TypeError("range expects one to three arguments") + if isinstance(step, _python.int) and step == 0: + raise ValueError("range step cannot be zero") + return _T.serial(start, stop, step=step) + + +__tvm_call_overrides__ = {_python.range: range_} + + +def logical_and(*values): + """Construct scalar/vector conjunction, preserving ordinary Python values.""" + if not values: + raise TypeError("logical_and requires at least one operand") + result = values[0] + for value in values[1:]: + if not isinstance(result, _ir.Expr) and not isinstance(value, _ir.Expr): + result = result and value + else: + lhs, rhs = _as_expr(result), _as_expr(value) + result = _tir.And(lhs, rhs) if lhs.ty.is_scalar() and rhs.ty.is_scalar() else lhs & rhs + return result + + +def logical_or(*values): + """Construct scalar/vector disjunction, preserving ordinary Python values.""" + if not values: + raise TypeError("logical_or requires at least one operand") + result = values[0] + for value in values[1:]: + if not isinstance(result, _ir.Expr) and not isinstance(value, _ir.Expr): + result = result or value + else: + lhs, rhs = _as_expr(result), _as_expr(value) + result = _tir.Or(lhs, rhs) if lhs.ty.is_scalar() and rhs.ty.is_scalar() else lhs | rhs + return result + + +def logical_not(value): + """Construct IR negation without coercing an IR expression to Python bool.""" + return _tir.Not(value) if isinstance(value, _ir.Expr) else not value + + +def select(condition, true_value, false_value): + """Construct a conditional expression or select an ordinary Python value.""" + if not isinstance(condition, _ir.Expr): + return true_value if condition else false_value + return _tir.Select(condition, true_value, false_value) diff --git a/python/tvm/tirx/script/v2.py b/python/tvm/tirx/script/v2.py new file mode 100644 index 0000000000..2b1668f082 --- /dev/null +++ b/python/tvm/tirx/script/v2.py @@ -0,0 +1,44 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +"""Opt-in TVMScript entry point using concrete TIRx construction operations.""" + +# pylint: disable=wildcard-import,unused-wildcard-import,redefined-builtin +import sys as _sys + +from tvm.script.parser_v2.frontend import make_decorator as _make_decorator +from tvm.script.parser_v2.frontend import make_helper as _make_helper +from tvm.script.parser_v2.frontend import register_namespace as _register_namespace + +from . import builder_v2 as _builder +from . import tile as _tile +from .builder_v2 import * # noqa: F403 +from .tile import cluster as cluster +from .tile import cta as cta +from .tile import thread as thread +from .tile import warp as warp +from .tile import warpgroup as warpgroup +from .tile import wg as wg + +tile = _tile +prim_func = _make_decorator( + _builder, option_map={"private": "private", "s_tir": "s_tir", "persistent": "persistent"} +) +inline = _make_helper(_builder, preserve_return=True) +macro = _make_helper(_builder, preserve_return=False) + +_register_namespace("T", _sys.modules[__name__]) +_register_namespace("tirx", _sys.modules[__name__])
