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/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..65e384e930c4 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) { @@ -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 { 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..6a90fa3c3ef0 --- /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-fA-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(