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 75ec20dc889256c9c45b768efbfa75b7cc7fd886 Author: Tianqi Chen <[email protected]> AuthorDate: Mon Sep 21 03:35:54 2026 +0000 Preserve native scalar storage semantics in concrete TIRx builders --- python/tvm/tirx/script/builder_v2/__init__.py | 65 ++++++++++++++------------- 1 file changed, 33 insertions(+), 32 deletions(-) diff --git a/python/tvm/tirx/script/builder_v2/__init__.py b/python/tvm/tirx/script/builder_v2/__init__.py index 1384cab01c..5903fc7374 100644 --- a/python/tvm/tirx/script/builder_v2/__init__.py +++ b/python/tvm/tirx/script/builder_v2/__init__.py @@ -208,25 +208,6 @@ def _as_expr(value): return _tir.const(value) -def _check_unterminated(): - frames = _IRBuilder.current().frames - if not frames or not isinstance(frames[-1], _frame.TIRFrame): - return - statements = frames[-1].stmts - while statements: - last = statements[-1] - if isinstance(last, _tir.Return | _tir.Break | _tir.Continue) or ( - isinstance(last, _tir.Evaluate) - and isinstance(last.value, _ir.Call) - and isinstance(last.value.op, _ir.Op) - and last.value.op.name in ("tirx.break_loop", "tirx.continue_loop") - ): - raise ValueError("An operation cannot follow an unconditional terminator") - if not isinstance(last, _tir.SeqStmt): - break - statements = last.seq - - def bind_( value=_MISSING, *, @@ -240,7 +221,6 @@ def bind_( ): """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, _frame.SBlockFrame): @@ -259,6 +239,15 @@ def bind_( elif isinstance(value, _ir.TensorLoad) and _tir.is_buffer_var(value.source): _name(value.source, name, name_span) return value + if previous is not _MISSING and _tir.is_buffer_var(previous): + shape = previous.ty.shape + if len(shape) == 1 and bool(shape[0] == 1): + if value is _MISSING: + raise ValueError("A reassignment requires an initializer") + _T.buffer_store(previous, value, [0]) + return previous + if previous is not _MISSING and isinstance(getattr(previous, "ty", None), _ir.PointerType): + raise ValueError(f"Pointer variable {name!r} cannot be reassigned") if previous is not _MISSING and ( _tir.is_buffer_var(previous) or isinstance(previous, _tir.IterVar) @@ -287,7 +276,7 @@ def bind_( if previous is not _MISSING and isinstance(previous, _ir.TensorLoad): if value is _MISSING: raise ValueError("A reassignment requires an initializer") - _T.buffer_store(previous.source, value, previous.indices) + _T.buffer_store(previous.source, value, list(previous.indices)) return previous if isinstance(value, _I.meta_var): return value.value @@ -302,8 +291,6 @@ def bind_( variable = _name(ty.as_var(rhs_dtype=value.ty), name, name_span) _T.Bind(value, var=variable) return variable - if value is _MISSING: - raise ValueError("Uninitialized scalar bindings are not supported") if ty is not None: annotation = ty() if callable(ty) and not isinstance(ty, _ir.Expr) else ty annotation = annotation.ty if isinstance(annotation, _ir.Expr) else annotation @@ -311,8 +298,11 @@ def bind_( raise TypeError("Mutable scalar annotations require a primitive scalar type") result = _T.local_scalar(str(annotation)).scalar _name(result.source, name, name_span) - _T.buffer_store(result.source, value, [0]) + if value is not _MISSING: + _T.buffer_store(result.source, value, [0]) return result + if value is _MISSING: + raise ValueError("An uninitialized binding requires a scalar type annotation") if ( isinstance(value, _ir.TensorLoad) and _tir.is_buffer_var(value.source) @@ -334,7 +324,7 @@ def bind_( return value if getattr(type(value), "_is_meta_class", False): if name is not None: - _T.name_meta_class_value(name, value) + _T.ir.name_meta_class_value(name, value) return value if _tir.is_buffer_var(value) or isinstance(value, _tir.IterVar | _tir.Layout): return _name(value, name, name_span) @@ -357,7 +347,6 @@ def emit_(value, *, span=None): """Consume one expression statement, including effect-only calls.""" if value is None or isinstance(value, str | _ir.Var): return - _check_unterminated() with _span_context(span): if isinstance(value, _NativeFrame | _Frame): _enter_concise(value) @@ -372,14 +361,29 @@ def emit_(value, *, span=None): def setitem(target, key, value, *, span=None): """Construct an indexed store after the caller has evaluated its operands.""" - _check_unterminated() with _span_context(span): _T.buffer_store(target, value, key) +def setattr(target, name, value, *, span=None): + """Store through scalar attributes, or update ordinary Python metadata.""" + if isinstance(value, _I.meta_var): + _python.setattr(target, name, value.value) + return + previous = getattr(target, name, _MISSING) + if isinstance(previous, _T.scalar_wrapper): + previous = previous.scalar + buffer = previous.source if isinstance(previous, _ir.TensorLoad) else previous + if _tir.is_buffer_var(buffer): + shape = buffer.ty.shape + if len(shape) == 1 and bool(shape[0] == 1): + bind_(value, previous=previous, span=span) + return + _python.setattr(target, name, value) + + def return_(value, *, span=None): - """Construct an IR return without exiting the Python construction helper.""" - _check_unterminated() + """Emit a native return; subsequent unreachable statements remain in the IR.""" if value is None: raise TypeError("A primitive function return requires an expression") with _span_context(span): @@ -398,7 +402,6 @@ def _require_loop(): def break_(*, span=None): """Construct a break targeting the nearest primitive loop.""" _require_loop() - _check_unterminated() with _span_context(span): _T.evaluate(_T.break_loop()) @@ -406,14 +409,12 @@ def break_(*, span=None): def continue_(*, span=None): """Construct a continue targeting the nearest primitive loop.""" _require_loop() - _check_unterminated() with _span_context(span): _T.evaluate(_T.continue_loop()) def assert_(condition, message="", *, span=None): """Emit the native flat assertion with its own source range.""" - _check_unterminated() kind = "RuntimeError" if isinstance(message, tuple): if len(message) != 2 or not isinstance(message[0], str):
