From 8174727bf16e1bfc4c68449c2ab9b16bef3ee2af Mon Sep 17 00:00:00 2001 From: Tianqi Chen Date: Fri, 2 Oct 2026 14:08:22 +0000 Subject: [PATCH 1/3] [FIX][Script] Preserve device identity and unify rendering configuration Use canonical device selectors and namespace aliases across translation and rendering. Keep range and tensor formatting in their owning hooks, avoid redundant nested imports and path mapping, and consolidate script entry points behind the public root API. --- include/tvm/script/printer/printer.h | 8 -- src/relax/script/printer/binding.cc | 2 - src/relax/script/printer/dependent_type.cc | 21 ++--- src/relax/script/printer/distributed.cc | 4 +- src/relax/script/printer/expr.cc | 7 +- src/relax/script/printer/function.cc | 6 +- src/relax/script/printer/global_info.cc | 3 +- src/relax/script/printer/utils.h | 2 - src/script/printer/doc_printer.cc | 4 +- src/script/printer/ir/ir.cc | 10 +-- src/script/printer/printer.cc | 67 +++++++++++++-- src/script/printer/script_printer.cc | 51 +----------- .../relax/script/test_relax_script_printer.py | 83 ++++++++++++++++--- .../script/test_script_printer_entry.py | 78 +++++++++++++++++ .../tirx/script/test_tirx_script_printer.py | 33 +++++++- 15 files changed, 256 insertions(+), 123 deletions(-) create mode 100644 tests/python/script/test_script_printer_entry.py diff --git a/include/tvm/script/printer/printer.h b/include/tvm/script/printer/printer.h index 22d4c52e565a..308d424bfaa1 100644 --- a/include/tvm/script/printer/printer.h +++ b/include/tvm/script/printer/printer.h @@ -54,14 +54,6 @@ TVM_DLL void RegisterNamespaceAlias(const ffi::String& key, const ffi::String& d /*! \brief Read the registered namespace aliases. */ TVM_DLL const ffi::Map& GetNamespaceAliases(); -/*! - * \brief Translate IR, recover diagnostic paths, and render Python text. - * \param obj The input IR object. - * \param config The translation and rendering options. - * \return The rendered script. - */ -TVM_DLL ffi::String Script(const ffi::ObjectRef& obj, const PrinterConfig& config); - } // namespace printer } // namespace script } // namespace tvm diff --git a/src/relax/script/printer/binding.cc b/src/relax/script/printer/binding.cc index 1eb978080f5a..5db6e124de25 100644 --- a/src/relax/script/printer/binding.cc +++ b/src/relax/script/printer/binding.cc @@ -64,8 +64,6 @@ ffi::Optional VarBindingDocTranslate(DocTranslatorObj* d, ffi::AnyView TVM_FFI_CHECK(destination == nullptr, TypeError) << "printer statement-only node cannot fulfill a destination"; if (auto func = binding->value.as()) { - d->Emit(CommentDoc("from tvm.script import relax as R"), ffi::GetRef(binding)); - d->Emit(CommentDoc(""), ffi::GetRef(binding)); IdDoc lhs = VarDoc(d, binding->var); d->Translate(binding->value); FunctionDoc function = d->CurrentScopeDocs().back().as_or_throw(); diff --git a/src/relax/script/printer/dependent_type.cc b/src/relax/script/printer/dependent_type.cc index 1dea84cb4ed4..e4a81d7a7c63 100644 --- a/src/relax/script/printer/dependent_type.cc +++ b/src/relax/script/printer/dependent_type.cc @@ -58,10 +58,11 @@ TVM_FFI_STATIC_INIT_BLOCK() { kDocTranslate, FDocTranslate::FromNative<&ShapeTypeDocTranslate>()); } -} // namespace - -ExprDoc RelaxTensorTypeDoc(DocTranslatorObj* d, const relax::TensorTypeNode* ty, - bool include_vdevice) { +ffi::Optional TensorTypeDocTranslate(DocTranslatorObj* d, ffi::AnyView input, + const ffi::Object*) { + const auto* ty = + ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck( + input); ffi::Array args; ffi::Array keys; ffi::Array values; @@ -82,7 +83,7 @@ ExprDoc RelaxTensorTypeDoc(DocTranslatorObj* d, const relax::TensorTypeNode* ty, keys.push_back("ndim"); values.push_back(LiteralDoc::Int(ty->ndim, std::nullopt)); } - if (include_vdevice && ty->vdevice.has_value()) { + if (ty->vdevice.has_value()) { keys.push_back("vdevice"); if (auto selector = GlobalInfoSelector(d, ty->vdevice.value())) { values.push_back(LiteralDoc::Str(selector.value(), std::nullopt)); @@ -94,16 +95,6 @@ ExprDoc RelaxTensorTypeDoc(DocTranslatorObj* d, const relax::TensorTypeNode* ty, return NamespaceDoc("relax")->Attr("Tensor")->Call(args, keys, values); } -namespace { - -ffi::Optional TensorTypeDocTranslate(DocTranslatorObj* d, ffi::AnyView input, - const ffi::Object*) { - const auto* ty = - ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck( - input); - return RelaxTensorTypeDoc(d, ty, true); -} - TVM_FFI_STATIC_INIT_BLOCK() { ffi::reflection::TypeAttrDef().attr( kDocTranslate, FDocTranslate::FromNative<&TensorTypeDocTranslate>()); diff --git a/src/relax/script/printer/distributed.cc b/src/relax/script/printer/distributed.cc index e0386ef8f764..97eb8a2e7f48 100644 --- a/src/relax/script/printer/distributed.cc +++ b/src/relax/script/printer/distributed.cc @@ -107,9 +107,7 @@ ffi::Optional DeviceMeshDocTranslate(DocTranslatorObj* d, ffi::AnyView } ExprDoc devices = LiteralDoc::None(std::nullopt); if (mesh->device_range.has_value()) { - CallDoc range = d->Translate(mesh->device_range.value()).value().as_or_throw(); - range->callee = NamespaceDoc("relax")->Attr("Range"); - devices = range; + devices = d->Translate(mesh->device_range.value()).value(); } else { ffi::Array ids; for (int64_t value : mesh->device_ids) ids.push_back(LiteralDoc::Int(value, std::nullopt)); diff --git a/src/relax/script/printer/expr.cc b/src/relax/script/printer/expr.cc index abd5180f4dd0..464cb69072ac 100644 --- a/src/relax/script/printer/expr.cc +++ b/src/relax/script/printer/expr.cc @@ -71,14 +71,9 @@ TVM_FFI_STATIC_INIT_BLOCK() { kDocTranslate, FDocTranslate::FromNative<&ShapeExprDocTranslate>()); } -ffi::Optional DataflowVarDocTranslate(DocTranslatorObj* d, ffi::AnyView input, - const ffi::Object* destination) { - return VarDocTranslate(d, input, destination); -} - TVM_FFI_STATIC_INIT_BLOCK() { ffi::reflection::TypeAttrDef().attr( - kDocTranslate, FDocTranslate::FromNative<&DataflowVarDocTranslate>()); + kDocTranslate, FDocTranslate::FromNative<&VarDocTranslate>()); } } // namespace diff --git a/src/relax/script/printer/function.cc b/src/relax/script/printer/function.cc index 42cba66e0acf..3e7eceb2f2ae 100644 --- a/src/relax/script/printer/function.cc +++ b/src/relax/script/printer/function.cc @@ -63,11 +63,7 @@ ffi::Optional FunctionDocTranslate(DocTranslatorObj* d, ffi::AnyView in auto signature_candidates = CopyImplicitDefs(d); ffi::Optional ret_type = std::nullopt; if (!func->ret_ty.IsMissing()) { - if (auto tensor = func->ret_ty.as()) { - ret_type = RelaxTensorTypeDoc(d, tensor, true); - } else { - ret_type = d->Translate(func->ret_ty).value(); - } + ret_type = d->Translate(func->ret_ty).value(); } ffi::Array decorator_keys; ffi::Array decorator_values; diff --git a/src/relax/script/printer/global_info.cc b/src/relax/script/printer/global_info.cc index d6ad01bf256d..af00725cd3b9 100644 --- a/src/relax/script/printer/global_info.cc +++ b/src/relax/script/printer/global_info.cc @@ -44,8 +44,7 @@ ffi::Optional GlobalInfoSelector(DocTranslatorObj* d, const GlobalI if (auto device = entry.as(); device && device.value()->target->kind->name == kind) { if (entry.same_as(info)) { - return ffi::String(std::string(kind) + ":" + std::to_string(index) + ":" + - std::string(device.value()->memory_scope)); + return ffi::String(std::string(kind) + ":" + std::to_string(index)); } ++index; } diff --git a/src/relax/script/printer/utils.h b/src/relax/script/printer/utils.h index 7cfba8b01a06..171b52c1b458 100644 --- a/src/relax/script/printer/utils.h +++ b/src/relax/script/printer/utils.h @@ -34,8 +34,6 @@ namespace details { ffi::Optional GlobalInfoSelector(DocTranslatorObj* d, const GlobalInfo& info); ExprDoc RelaxShapeDim(DocTranslatorObj* d, const PrimExpr& dim); -ExprDoc RelaxTensorTypeDoc(DocTranslatorObj* d, const relax::TensorTypeNode* ty, - bool include_vdevice); ffi::Array RelaxSeqBody(DocTranslatorObj* d, const relax::SeqExprNode* seq, ffi::Optional destination = std::nullopt, ffi::Optional annotation = std::nullopt, diff --git a/src/script/printer/doc_printer.cc b/src/script/printer/doc_printer.cc index 37080b6c4d59..bfb6cbecf14d 100644 --- a/src/script/printer/doc_printer.cc +++ b/src/script/printer/doc_printer.cc @@ -1159,8 +1159,8 @@ void PythonDocPrinter::PrintTypedDoc(const NamespaceDoc& doc) { const ffi::String& name = doc->canonical_name; ffi::String key = std::string(name) + ".prefix"; ffi::String fallback = GetNamespaceAliases().Get(key).value_or(name); - output_ << (name == "ir" ? config()->ir_prefix - : config()->GetExtraConfig(key, fallback)); + output_ << config()->GetExtraConfig(key, + name == "ir" ? config()->ir_prefix : fallback); } void PythonDocPrinter::PrintTypedDoc(const AttrAccessDoc& doc) { diff --git a/src/script/printer/ir/ir.cc b/src/script/printer/ir/ir.cc index 43ec20006213..29c3c4d9fba2 100644 --- a/src/script/printer/ir/ir.cc +++ b/src/script/printer/ir/ir.cc @@ -84,15 +84,7 @@ ffi::Optional IRModuleDocTranslate(DocTranslatorObj* d, ffi::AnyView in for (const auto& [key, entries] : infos) { ffi::Array items; for (const GlobalInfo& entry : entries) { - ExprDoc item = AnyValue(d, entry); - if (key == "mesh") { - if (auto mesh = item.as(); mesh && mesh.value()->args.size() == 2) { - if (auto range = mesh.value()->args[1].as()) { - range.value()->callee = NamespaceDoc("ir")->Attr("Range"); - } - } - } - items.push_back(item); + items.push_back(AnyValue(d, entry)); } keys.push_back(LiteralDoc::Str(key, std::nullopt)); values.push_back(ListDoc(items)); diff --git a/src/script/printer/printer.cc b/src/script/printer/printer.cc index a8ca135d9f1f..8387d6fb0d45 100644 --- a/src/script/printer/printer.cc +++ b/src/script/printer/printer.cc @@ -20,12 +20,14 @@ #include #include #include +#include #include #include #include #include #include +#include #include #include #include @@ -108,8 +110,8 @@ ffi::Map> CollectMetadata( ffi::Array DisplayAliases(const PrinterConfig& config) { ffi::Array aliases; for (const auto& [key, fallback] : GetNamespaceAliases()) { - aliases.push_back(key == "ir.prefix" ? config->ir_prefix - : config->GetExtraConfig(key, fallback)); + aliases.push_back(config->GetExtraConfig( + key, key == "ir.prefix" ? config->ir_prefix : fallback)); } return aliases; } @@ -330,9 +332,7 @@ class MapDocPaths { std::vector> occurrences_; }; -} // namespace - -ffi::String Script(const ffi::ObjectRef& obj, const PrinterConfig& config) { +ffi::String RenderScript(const ffi::ObjectRef& obj, const PrinterConfig& config) { ffi::Dict origins; Doc doc = DocTranslate(obj, &origins, config->extra_config); auto block = doc.as(); @@ -410,8 +410,8 @@ ffi::String Script(const ffi::ObjectRef& obj, const PrinterConfig& config) { std::string canonical = key; canonical.resize(canonical.size() - std::string(".prefix").size()); if (!namespaces.count(canonical)) continue; - ffi::String alias = canonical == "ir" ? config->ir_prefix - : config->GetExtraConfig(key, fallback); + ffi::String alias = config->GetExtraConfig( + key, canonical == "ir" ? config->ir_prefix : fallback); ffi::String import = "from tvm.script import " + canonical + " as " + std::string(alias); if (config->GetExtraConfig("ir.comment_imports", false) && !config->show_meta) { header.push_back(CommentDoc(import)); @@ -421,13 +421,64 @@ ffi::String Script(const ffi::ObjectRef& obj, const PrinterConfig& config) { } if (!header.empty()) header.push_back(ffi::String("\n")); } + if (config->path_to_underline.empty() && config->obj_to_underline.empty() && + config->path_to_annotate.empty() && config->obj_to_annotate.empty() && + !config->render_invisible_path_info) { + return details::RenderPythonScript(doc, config, header, {}, {}); + } MapDocPaths paths(obj, config); auto [underline_paths, annotations] = paths.Map(doc, origins); return details::RenderPythonScript(doc, config, header, underline_paths, annotations); } -TVM_FFI_STATIC_INIT_BLOCK() { ffi::reflection::GlobalDef().def("script.printer.Script", Script); } +} // namespace } // namespace printer } // namespace script +namespace { + +std::string RenderFallbackWithInvisiblePathInfo(const ffi::String& script, + const PrinterConfig& config) { + if (!config->render_invisible_path_info || config->path_to_underline.empty()) { + return std::string(script); + } + + std::ostringstream os; + for (size_t i = 0; i < config->path_to_underline.size(); ++i) { + if (i != 0) os << "\n"; + os << "Access path: " << config->path_to_underline[i] + << "\nNote: No visible object for this path is rendered in TVMScript."; + } + os << "\n\n" << script; + return os.str(); +} + +} // namespace + +std::string Script(const ffi::ObjectRef& node, const ffi::Optional& cfg) { + PrinterConfig config = cfg.value_or(PrinterConfig()); + static ffi::reflection::TypeAttrColumn translate(script::printer::kDocTranslate); + // Builtin runtime roots keep their native repr; hooks still translate them within IR. + if (!node.defined() || node->type_index() < ffi::TypeIndex::kTVMFFIDynObjectBegin || + translate[node->type_index()].type_index() == ffi::TypeIndex::kTVMFFINone) { + return RenderFallbackWithInvisiblePathInfo(ffi::ReprPrint(ffi::Any(node)), config); + } + return std::string(script::printer::RenderScript(node, config)); +} + +std::string RedirectedReprPrinterMethod(const ffi::ObjectRef& obj) { + try { + PrinterConfig config; + config->extra_config.Set("ir.comment_imports", true); + // Call translation directly so an unsupported type cannot recurse through ffi repr. + return std::string(script::printer::RenderScript(obj, config)); + } catch (const tvm::ffi::Error& e) { + LOG(WARNING) << "TVMScript printer falls back to the basic address printer with the error:\n" + << e.what(); + std::ostringstream os; + os << obj->GetTypeKey() << '(' << obj.get() << ')'; + return os.str(); + } +} + } // namespace tvm diff --git a/src/script/printer/script_printer.cc b/src/script/printer/script_printer.cc index 6a712759318f..4c6ce2337c0a 100644 --- a/src/script/printer/script_printer.cc +++ b/src/script/printer/script_printer.cc @@ -27,8 +27,6 @@ #include #include #include -#include -#include #include #include #include @@ -38,66 +36,21 @@ #include #include -#include #include #include "utils.h" namespace tvm { -namespace { - -std::string RenderFallbackWithInvisiblePathInfo(const ffi::String& script, - const PrinterConfig& config) { - if (!config->render_invisible_path_info || config->path_to_underline.empty()) { - return std::string(script); - } - - std::ostringstream os; - for (size_t i = 0; i < config->path_to_underline.size(); ++i) { - if (i != 0) os << "\n"; - os << "Access path: " << config->path_to_underline[i] - << "\nNote: No visible object for this path is rendered in TVMScript."; - } - os << "\n\n" << script; - return os.str(); -} - -} // namespace - -std::string Script(const ffi::ObjectRef& node, const ffi::Optional& cfg) { - PrinterConfig config = cfg.value_or(PrinterConfig()); - static ffi::reflection::TypeAttrColumn translate(script::printer::kDocTranslate); - // Builtin runtime roots keep their native repr; hooks still translate them within IR. - if (!node.defined() || node->type_index() < ffi::TypeIndex::kTVMFFIDynObjectBegin || - translate[node->type_index()].type_index() == ffi::TypeIndex::kTVMFFINone) { - return RenderFallbackWithInvisiblePathInfo(ffi::ReprPrint(ffi::Any(node)), config); - } - return std::string(script::printer::Script(node, config)); -} - -std::string RedirectedReprPrinterMethod(const ffi::ObjectRef& obj) { - try { - PrinterConfig config; - config->extra_config.Set("ir.comment_imports", true); - // Call translation directly so an unsupported type cannot recurse through ffi repr. - return std::string(script::printer::Script(obj, config)); - } catch (const tvm::ffi::Error& e) { - LOG(WARNING) << "TVMScript printer falls back to the basic address printer with the error:\n" - << e.what(); - std::ostringstream os; - os << obj->GetTypeKey() << '(' << obj.get() << ')'; - return os.str(); - } -} TVM_FFI_STATIC_INIT_BLOCK() { namespace refl = ffi::reflection; using script::printer::details::RegisterScriptRepr; refl::GlobalDef() .def("node.TVMScriptPrinterScript", tvm::Script) + .def("script.printer.Script", tvm::Script) .def("script.printer.ReprPrintRelax", [](const ffi::ObjectRef& obj, const PrinterConfig& config) { - return script::printer::Script(obj, config); + return tvm::Script(obj, config); }); RegisterScriptRepr(); diff --git a/tests/python/relax/script/test_relax_script_printer.py b/tests/python/relax/script/test_relax_script_printer.py index 1d3a507127b5..4e6e391a6d1c 100644 --- a/tests/python/relax/script/test_relax_script_printer.py +++ b/tests/python/relax/script/test_relax_script_printer.py @@ -49,7 +49,7 @@ def test_constant(): ) assert ( constant.__str__() - == """R.dist.const(1.0, R.DTensor((), "float32", R.device_mesh((2, 2), R.Range(0, 4)), "R, R"))""" + == """R.dist.const(1.0, R.DTensor((), "float32", R.device_mesh((2, 2), I.Range(0, 4)), "R, R"))""" ) @@ -59,13 +59,13 @@ def test_dtensor_type(): obj0 = DTensorType(tensor_ty1, DeviceMesh((2, 2), Range(0, 4)), Placement.from_text("S[1], R")) assert ( obj0.__str__() - == """R.DTensor((32, 32), "float32", R.device_mesh((2, 2), R.Range(0, 4)), "S[1], R")""" + == """R.DTensor((32, 32), "float32", R.device_mesh((2, 2), I.Range(0, 4)), "S[1], R")""" ) obj1 = DTensorType(tensor_ty2, DeviceMesh((2, 2), Range(0, 4)), Placement.from_text("S[1], R")) assert ( obj1.__str__() - == """R.DTensor((32, 32), device_mesh=R.device_mesh((2, 2), R.Range(0, 4)), placement="S[1], R")""" + == """R.DTensor((32, 32), device_mesh=R.device_mesh((2, 2), I.Range(0, 4)), placement="S[1], R")""" ) obj2 = DTensorType(tensor_ty2, DeviceMesh((2, 2), [0, 1, 2, 3]), Placement.from_text("S[1], R")) @@ -75,6 +75,45 @@ def test_dtensor_type(): ) +@pytest.mark.parametrize("memory_scope", ["", "global", "global:workspace"]) +def test_vdevice_selector_preserves_registry_identity(memory_scope): + @I.ir_module + class Module: + I.module_global_infos( + { + "vdevice": [ + R.vdevice("llvm"), + R.vdevice("cuda", 0, "global"), + R.vdevice("cuda", 0, memory_scope), + ] + } + ) + + @R.function + def main(x: R.Tensor((4,), "float32", "cuda:1")) -> R.Tensor((4,), "float32", "cuda:1"): + return x + + script = Module.script() + assert script.count('vdevice="cuda:1"') == 2 + restored = tvm.script.from_source(script, extra_vars={"I": I, "R": R, "T": T}) + tvm.ir.assert_structural_equal(Module, restored) + selected = restored.global_infos["vdevice"][2] + assert selected.memory_scope == memory_scope + assert restored["main"].params[0].ty.vdevice.is_(selected) + assert restored["main"].ret_ty.vdevice.is_(selected) + assert not selected.is_(restored.global_infos["vdevice"][1]) + + +def test_device_mesh_range_namespace_roundtrip(): + mesh = DeviceMesh((2, 2), Range(0, 4)) + assert str(mesh) == "R.device_mesh((2, 2), I.Range(0, 4))" + mod = IRModule({}) + mod.update_global_info("mesh", [mesh]) + script = mod.script() + assert "R.device_mesh((2, 2), I.Range(0, 4))" in script + tvm.ir.assert_structural_equal(mod, tvm.script.from_source(script, extra_vars={"I": I, "R": R})) + + @I.ir_module class TestModule: I.module_attrs({"device_num": 10}) @@ -118,11 +157,12 @@ def test_func(): """ from __future__ import annotations +# from tvm.script import ir as I # from tvm.script import relax as R @R.function -def foo(x: R.DTensor((128, 128), "float32", R.device_mesh((2, 2), R.Range(0, 4)), "S[0], R")) -> R.DTensor((128, 128), "float32", R.device_mesh((2, 2), R.Range(0, 4)), "S[0], R"): - gv0 = R.dist.call_tir(Module.tir_func, (x,), out_ty=R.DTensor((128, 128), "float32", R.device_mesh((2, 2), R.Range(0, 4)), "S[0], R")) +def foo(x: R.DTensor((128, 128), "float32", R.device_mesh((2, 2), I.Range(0, 4)), "S[0], R")) -> R.DTensor((128, 128), "float32", R.device_mesh((2, 2), I.Range(0, 4)), "S[0], R"): + gv0 = R.dist.call_tir(Module.tir_func, (x,), out_ty=R.DTensor((128, 128), "float32", R.device_mesh((2, 2), I.Range(0, 4)), "S[0], R")) return gv0 """, ) @@ -195,6 +235,18 @@ def func(a: R.Tensor((10, 10))) -> R.Tensor((10, 10)): ) +def test_function_return_tensor_source_span(): + x = relax.Var("x", relax.TensorType((4,), "float32")) + func = relax.Function([x], x, ret_ty=x.ty).with_attr("global_symbol", "main") + script = func.script(path_to_underline=[AccessPath.root().attr("ret_ty")]) + lines = script.splitlines() + index = next(i for i, line in enumerate(lines) if "def main" in line) + definition, underline = lines[index : index + 2] + start = definition.index(" -> ") + len(" -> ") + end = definition.rindex(":") + assert underline == " " * start + "^" * (end - start) + + def test_function_dependent_shape_source_spans(): n = tirx.Var("n", "int64") cast = tirx.Cast("int64", n) @@ -336,7 +388,9 @@ def test_extern_func_with_ty_roundtrip(): tvm.ir.assert_structural_equal(mod, roundtrip) -def test_nested_function(): +@pytest.mark.parametrize("relax_prefix", ["R", "CustomR"]) +@pytest.mark.parametrize("comment_imports", [False, True]) +def test_nested_function(relax_prefix, comment_imports): @I.ir_module class NestedFunction: @R.function @@ -348,8 +402,12 @@ def nested(y: R.Tensor((), "int32")) -> R.Tensor((), "int32"): z = nested(x) return z - _assert_print_lines( - NestedFunction, + script = NestedFunction.script( + verbose_expr=True, + extra_config={"relax.prefix": relax_prefix, "ir.comment_imports": comment_imports}, + ) + _assert_print( + script, """ from __future__ import annotations @@ -360,15 +418,18 @@ def nested(y: R.Tensor((), "int32")) -> R.Tensor((), "int32"): class Module: @R.function def main(x: R.Tensor((), dtype="int32")) -> R.Tensor((), dtype="int32"): - # from tvm.script import relax as R - @R.function def nested(y: R.Tensor((), dtype="int32")) -> R.Tensor((), dtype="int32"): return y z: R.Tensor((), dtype="int32") = nested(x) return z -""", +""".replace("R.", f"{relax_prefix}.") + .replace("relax as R", f"relax as {relax_prefix}") + .replace("# from", "# from" if comment_imports else "from"), + ) + tvm.ir.assert_structural_equal( + NestedFunction, tvm.script.from_source(script, extra_vars={"I": I, relax_prefix: R}) ) diff --git a/tests/python/script/test_script_printer_entry.py b/tests/python/script/test_script_printer_entry.py new file mode 100644 index 000000000000..7fe5e2ef349e --- /dev/null +++ b/tests/python/script/test_script_printer_entry.py @@ -0,0 +1,78 @@ +# 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. +"""Public script entry points and diagnostic rendering.""" + +import re + +import pytest +import tvm_ffi +from tvm_ffi.access_path import AccessPath + +import tvm +from tvm.runtime.script_printer import PrinterConfig, _script + + +@pytest.mark.parametrize( + "entry", + ["node.TVMScriptPrinterScript", "script.printer.Script", "script.printer.ReprPrintRelax"], +) +def test_script_entry_points(entry): + render = tvm_ffi.get_global_func(entry) + config = PrinterConfig(extra_config={"ir.prefix": "IR"}) + assert render(tvm.ir.Range(0, 4), config) == "IR.Range(0, 4)" + fallback_config = PrinterConfig(path_to_underline=[AccessPath.root().attr("missing")]) + assert render(tvm.runtime.ShapeTuple([1, 2]), fallback_config) == ( + "Access path: .missing\n" + "Note: No visible object for this path is rendered in TVMScript.\n\n" + "Shape(1, 2)" + ) + + +@pytest.mark.parametrize( + "diagnostic", ["path_to_underline", "obj_to_underline", "path_to_annotate", "obj_to_annotate"] +) +def test_diagnostics_without_invisible_path_info(diagnostic): + value = tvm.ir.PrimType("int32") + target = AccessPath.root() if diagnostic.startswith("path") else value + annotated = diagnostic.endswith("annotate") + config = PrinterConfig( + **{diagnostic: {target: "type note"} if annotated else [target]}, + extra_config={"render_invisible_path_info": False}, + ) + text = _script(value, config) + assert "T.int32" in text + assert "type note" in text if annotated else "^^^^^^^" in text + assert "Access path:" not in text + + +def test_plain_render_without_path_mapping(): + mod = tvm.IRModule(attrs={"constant": tvm.runtime.tensor([1, 2, 3])}) + options = {"ir.prefix": "IR", "ir.module_name": "Example"} + expected = mod.script(show_meta=True, extra_config=options) + actual = mod.script( + show_meta=True, extra_config={**options, "render_invisible_path_info": False} + ) + assert actual == expected + assert "load_json" in actual + assert "from tvm.script import ir as IR" in actual + + +def test_redirected_repr_translation_failure(): + value = tvm.tirx.For(tvm.tirx.Var("i", "int32"), 0, 1, 99, tvm.tirx.Evaluate(0)) + with pytest.raises(TypeError, match="unknown loop kind"): + value.script() + assert re.fullmatch(r"tirx\.For\(0x[0-9a-f]+\)", repr(value)) diff --git a/tests/python/tirx/script/test_tirx_script_printer.py b/tests/python/tirx/script/test_tirx_script_printer.py index e352d4368ee1..0f86ba147bf1 100644 --- a/tests/python/tirx/script/test_tirx_script_printer.py +++ b/tests/python/tirx/script/test_tirx_script_printer.py @@ -74,7 +74,7 @@ def test_config_extension_passthrough(): assert _script(tirx.Var("Custom", "int32"), cfg) == "Custom_1" -@pytest.mark.parametrize("key", ["tirx.prefix", "relax.prefix", "s_tir.prefix"]) +@pytest.mark.parametrize("key", ["ir.prefix", "tirx.prefix", "relax.prefix", "s_tir.prefix"]) @pytest.mark.parametrize("value", ["2prefix", 17]) @pytest.mark.parametrize("nested", [False, True]) def test_config_validates_dialect_prefixes(key, value, nested): @@ -99,6 +99,37 @@ def test_config_reserves_dialect_prefixes_before_variable_definition(prefixes): ) +@pytest.mark.parametrize( + "config, alias", + [ + ({}, "I"), + ({"ir_prefix": "TypedI"}, "TypedI"), + ({"extra_config": {"ir.prefix": "ExplicitI"}}, "ExplicitI"), + ( + {"ir_prefix": "TypedI", "extra_config": {"ir.prefix": "ExplicitI"}}, + "ExplicitI", + ), + ], +) +@pytest.mark.parametrize("comment_imports", [False, True]) +def test_ir_prefix_rendering_imports_and_reservations(config, alias, comment_imports): + config = dict(config) + config["extra_config"] = { + **config.get("extra_config", {}), + "ir.comment_imports": comment_imports, + } + var = tirx.Var(alias, "int32") + assert var.script(verbose_expr=True, **config).strip() == ( + f'{alias}_1 = {alias}.dynamic("{alias}", dtype="int32")\n{alias}_1' + ) + mod = tvm.IRModule({}) + script = mod.script(**config) + import_line = f"from tvm.script import ir as {alias}" + assert script.splitlines()[0] == ("# " if comment_imports else "") + import_line + assert f"@{alias}.ir_module" in script + assert_structural_equal(mod, tvm.script.from_source(script, extra_vars={alias: I})) + + def test_buffer(): a = tirx.decl_tensor((128, 128), "float16", name="A") _assert_print( From 6a089851f5184b8a4186deed5ae2a76304b4e855 Mon Sep 17 00:00:00 2001 From: Tianqi Chen Date: Fri, 2 Oct 2026 16:48:23 +0000 Subject: [PATCH 2/3] [REFACTOR][Script] Keep document rendering internal Remove redundant public C++ Doc rendering helpers and bind the Python document entry point directly to the internal renderer. --- docs/arch/tvmscript.rst | 9 +++++---- include/tvm/script/printer/doc.h | 19 ------------------- src/script/printer/doc_printer.cc | 16 +++++----------- src/script/printer/doc_printer.h | 1 + 4 files changed, 11 insertions(+), 34 deletions(-) diff --git a/docs/arch/tvmscript.rst b/docs/arch/tvmscript.rst index c7441abbc6ae..7b5feed65543 100644 --- a/docs/arch/tvmscript.rst +++ b/docs/arch/tvmscript.rst @@ -158,11 +158,12 @@ Printing and round trips of expressions and statements. The translation engine tracks scopes, names and each Doc's original IR object. The script entry point maps these origins to diagnostic paths; the private Doc printer formats the tree, annotations and underlines as Python -text. Printer configuration stays read-only throughout. ``DocToPythonScript`` also -formats an existing Doc directly. This tree is separate from the parser's Python AST. +text. Printer configuration stays read-only throughout. The Python document printer +helper also formats an existing Doc directly for document-level tests. This tree is +separate from the parser's Python AST. -The public ``Script`` text entry points are declared in -``tvm/script/printer/printer.h``. Their orchestration and diagnostic path mapping +The public ``tvm::Script`` text entry point is declared in +``tvm/script/printer/printer.h``. Its orchestration and diagnostic path mapping live in ``src/script/printer/printer.cc``; ``doc_translator.h`` exposes the IR-to-Doc translation protocol. diff --git a/include/tvm/script/printer/doc.h b/include/tvm/script/printer/doc.h index eba976d265ac..1b5c6283038d 100644 --- a/include/tvm/script/printer/doc.h +++ b/include/tvm/script/printer/doc.h @@ -25,7 +25,6 @@ #include #include #include -#include #include @@ -1435,24 +1434,6 @@ class NamespaceDoc : public ExprDoc { TVM_FFI_DEFINE_OBJECT_REF_METHODS_NOTNULLABLE(NamespaceDoc, ExprDoc, NamespaceDocNode); }; -/*! - * \brief Render a completed Doc tree with the supplied configuration. - * \param doc The Doc tree to render. - * \param config The rendering configuration. - * \return The rendered Python source. - */ -TVM_DLL ffi::String DocToPythonScript(Doc doc, const PrinterConfig& config); - -/*! - * \brief Render explicit headers followed by a completed Doc tree. - * \param doc The Doc tree to render. - * \param config The rendering configuration. - * \param header Ordered String source chunks or CommentDocs. - * \return The rendered Python source. - */ -TVM_DLL ffi::String DocToPythonScriptWithHeader(Doc doc, const PrinterConfig& config, - const ffi::Array& header); - } // namespace printer } // namespace script } // namespace tvm diff --git a/src/script/printer/doc_printer.cc b/src/script/printer/doc_printer.cc index bfb6cbecf14d..65e384e930c4 100644 --- a/src/script/printer/doc_printer.cc +++ b/src/script/printer/doc_printer.cc @@ -1698,7 +1698,7 @@ ffi::String RenderPythonScript(Doc doc, const PrinterConfig& cfg, std::string script = printer.GetString(); // GetString terminates non-empty output with one newline. Preserve the - // established DocToPythonScript result without normalizing any other + // established rendering result without normalizing any other // trailing whitespace. if (!script.empty()) { TVM_FFI_ICHECK_EQ(script.back(), '\n'); @@ -1711,18 +1711,12 @@ ffi::String RenderPythonScript(Doc doc, const PrinterConfig& cfg, } // namespace details -ffi::String DocToPythonScriptWithHeader(Doc doc, const PrinterConfig& cfg, - const ffi::Array& header) { - return details::RenderPythonScript(std::move(doc), cfg, header, cfg->path_to_underline, {}); -} - -ffi::String DocToPythonScript(Doc doc, const PrinterConfig& cfg) { - return DocToPythonScriptWithHeader(std::move(doc), cfg, {}); -} - TVM_FFI_STATIC_INIT_BLOCK() { namespace refl = tvm::ffi::reflection; - refl::GlobalDef().def("script.printer.DocToPythonScript", DocToPythonScript); + refl::GlobalDef().def("script.printer.DocToPythonScript", [](Doc doc, + const PrinterConfig& config) { + return details::RenderPythonScript(std::move(doc), config, {}, config->path_to_underline, {}); + }); } } // namespace printer diff --git a/src/script/printer/doc_printer.h b/src/script/printer/doc_printer.h index 819e3467db5c..f7a5a6a340ec 100644 --- a/src/script/printer/doc_printer.h +++ b/src/script/printer/doc_printer.h @@ -19,6 +19,7 @@ #ifndef SRC_SCRIPT_PRINTER_DOC_PRINTER_H_ #define SRC_SCRIPT_PRINTER_DOC_PRINTER_H_ +#include #include namespace tvm { From c8f8aa66bfad930e0d8b72e2fe22323752044af1 Mon Sep 17 00:00:00 2001 From: Tianqi Chen Date: Sun, 4 Oct 2026 02:14:03 +0000 Subject: [PATCH 3/3] [FIX][Script] Accept platform pointer formatting in repr assertion Allow the native hexadecimal pointer representation with or without a prefix and either letter case. --- tests/python/script/test_script_printer_entry.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/python/script/test_script_printer_entry.py b/tests/python/script/test_script_printer_entry.py index 7fe5e2ef349e..6a90fa3c3ef0 100644 --- a/tests/python/script/test_script_printer_entry.py +++ b/tests/python/script/test_script_printer_entry.py @@ -75,4 +75,4 @@ def test_redirected_repr_translation_failure(): value = tvm.tirx.For(tvm.tirx.Var("i", "int32"), 0, 1, 99, tvm.tirx.Evaluate(0)) with pytest.raises(TypeError, match="unknown loop kind"): value.script() - assert re.fullmatch(r"tirx\.For\(0x[0-9a-f]+\)", repr(value)) + assert re.fullmatch(r"tirx\.For\((?:0x)?[0-9a-fA-F]+\)", repr(value))