diff --git a/include/tvm/tirx/script/ir_builder/frame.h b/include/tvm/tirx/script/ir_builder/frame.h index 295a59870301..97f47f2aa53c 100644 --- a/include/tvm/tirx/script/ir_builder/frame.h +++ b/include/tvm/tirx/script/ir_builder/frame.h @@ -531,45 +531,6 @@ class DeclBufferFrame : public TIRFrame { TVM_FFI_DEFINE_OBJECT_REF_METHODS_NOTNULLABLE(DeclBufferFrame, TIRFrame, DeclBufferFrameNode); }; -/*! - * \brief A frame that represents a hint directive for the sketch language. - * - * \sa HintFrame - */ -class HintFrameNode : public TIRFrameNode { - public: - /*! \brief The free-form hint message string. */ - ffi::String message; - /*! \brief Optional structured key-value attributes. */ - ffi::Map attrs; - - static void RegisterReflection() { - namespace refl = tvm::ffi::reflection; - refl::ObjectDef() - .def_ro("message", &HintFrameNode::message) - .def_ro("attrs", &HintFrameNode::attrs); - } - TVM_FFI_DECLARE_OBJECT_INFO_FINAL("script.ir_builder.tirx.HintFrame", HintFrameNode, - TIRFrameNode); - - public: - void ExitWithScope() final; -}; - -/*! - * \brief Managed reference to HintFrameNode. - * - * \sa HintFrameNode - */ -class HintFrame : public TIRFrame { - public: - explicit HintFrame(ffi::ObjectPtr data) : TIRFrame(ffi::UnsafeInit{}) { - TVM_FFI_ICHECK(data != nullptr); - data_ = std::move(data); - } - TVM_FFI_DEFINE_OBJECT_REF_METHODS_NOTNULLABLE(HintFrame, TIRFrame, HintFrameNode); -}; - } // namespace tirx } // namespace ir_builder } // namespace script diff --git a/include/tvm/tirx/stmt.h b/include/tvm/tirx/stmt.h index 212ad12b980a..7abfb6621040 100644 --- a/include/tvm/tirx/stmt.h +++ b/include/tvm/tirx/stmt.h @@ -809,12 +809,6 @@ constexpr const char* device_id = "device_id"; constexpr const char* device_scope = "device_scope"; /*! \brief The device type. */ constexpr const char* device_type = "device_type"; -/*! - * \brief Mark the scope as generated by extern primitive. - * Such scope can contain arbitrary ir program and we need to be careful - * when making certain assumptions about the structure of the program. - */ -constexpr const char* extern_scope = "extern_scope"; /*! \brief Pragma: auto-unroll, max_step */ constexpr const char* pragma_auto_unroll_max_step = "pragma_auto_unroll_max_step"; /*! \brief Import C source or file into the final code gen module */ diff --git a/python/tvm/tirx/script/ir_builder/frame.py b/python/tvm/tirx/script/ir_builder/frame.py index 2aaedb4e71ff..c46a97af0bb7 100644 --- a/python/tvm/tirx/script/ir_builder/frame.py +++ b/python/tvm/tirx/script/ir_builder/frame.py @@ -100,7 +100,3 @@ class LaunchThreadFrame(TIRFrame): def __enter__(self) -> Var: super().__enter__() return self.iter_var.var - - -@_register_object("script.ir_builder.tirx.HintFrame") -class HintFrame(TIRFrame): ... diff --git a/python/tvm/tirx/script/ir_builder/parser_protocol.py b/python/tvm/tirx/script/ir_builder/parser_protocol.py index 90fa9184c729..268dabd861ff 100644 --- a/python/tvm/tirx/script/ir_builder/parser_protocol.py +++ b/python/tvm/tirx/script/ir_builder/parser_protocol.py @@ -503,24 +503,6 @@ def attr( return _ffi_api.Attr(node_or_dict, attr_key, value) # type: ignore[attr-defined] # pylint: disable=no-member -def hint(message: str = "", **attrs) -> frame.HintFrame: - """Universal directive primitive for the sketch language. - - Parameters - ---------- - message : str - Free-form directive string that the agent interprets. - **attrs - Optional structured key-value attributes for known patterns. - - Returns - ------- - res : frame.HintFrame - Usable as context manager (with T.hint("msg"):) or bare statement (T.hint("msg")). - """ - return _ffi_api.Hint(message, attrs or {}) # type: ignore[attr-defined] # pylint: disable=no-member - - def buffer_store( buffer: Buffer, # pylint: disable=redefined-outer-name value: Expr, @@ -1079,7 +1061,6 @@ def env_thread(thread_tag: str, dtype: str = "int32") -> Var: "ge_", "grid", "gt_", - "hint", "if_", "if_then_else_", "launch_thread", diff --git a/src/backend/cuda/codegen/codegen_cuda.cc b/src/backend/cuda/codegen/codegen_cuda.cc index 7cba4d4c1dfa..e71b272d71fa 100644 --- a/src/backend/cuda/codegen/codegen_cuda.cc +++ b/src/backend/cuda/codegen/codegen_cuda.cc @@ -1619,20 +1619,6 @@ void CodeGenCUDA::Dispatch_(const AttrStmtNode* op) { TVM_FFI_ICHECK(inner); this->Dispatch(inner->body); return; - } else if (op->attr_key == "disable_unroll") { - PrintIndent(); - stream << "#pragma unroll 1\n"; - this->Dispatch(op->body); - return; - } else if (op->attr_key == "pragma_unroll") { - PrintIndent(); - stream << "#pragma unroll"; - if (const auto* count = op->value.as(); count && count->value != 1) { - stream << " " << count->value; - } - stream << "\n"; - this->Dispatch(op->body); - return; } else if (op->attr_key == tirx::attr::thread_extent) { } CodeGenC::Dispatch_(op); diff --git a/src/s_tir/transform/merge_shared_memory_allocations.cc b/src/s_tir/transform/merge_shared_memory_allocations.cc index e9b2db0c0b3d..fd843692393d 100644 --- a/src/s_tir/transform/merge_shared_memory_allocations.cc +++ b/src/s_tir/transform/merge_shared_memory_allocations.cc @@ -295,8 +295,6 @@ class SharedMemLinearAccessPatternFinder final : public StmtExprVisitor { in_thread_env_ = true; TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(VisitNewScope(op)); in_thread_env_ = false; - } else if (op->attr_key == tirx::attr::extern_scope) { - TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(VisitNewScope(op)); } else if (op->attr_key == s_tir::attr::virtual_thread) { TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(VisitNewScope(op)); } else { @@ -461,7 +459,6 @@ class SharedMemoryRewriter : public StmtExprMutator { } in_thread_env_ = false; - // 6. If this scope has no shmem allocs, skip the wrapper. if (scope.shmem_allocs.empty()) { scope_stack_.pop_back(); diff --git a/src/tirx/script/ir_builder/frame.cc b/src/tirx/script/ir_builder/frame.cc index cc10fb0fea93..760fcfd539d1 100644 --- a/src/tirx/script/ir_builder/frame.cc +++ b/src/tirx/script/ir_builder/frame.cc @@ -49,7 +49,6 @@ TVM_FFI_STATIC_INIT_BLOCK() { ThenFrameNode::RegisterReflection(); ElseFrameNode::RegisterReflection(); DeclBufferFrameNode::RegisterReflection(); - HintFrameNode::RegisterReflection(); } namespace { @@ -291,21 +290,6 @@ void DeclBufferFrameNode::ExitWithScope() { } } -void HintFrameNode::ExitWithScope() { - TIRFrameNode::ExitWithScope(); - // Always store attrs as a structured Map in the node field - ffi::Map full_attrs; - if (!message.empty()) { - full_attrs.Set("message", ffi::String(message)); - } - for (const auto& [k, v] : attrs) { - full_attrs.Set(k, v); - } - AddToParent( - tvm::tirx::AttrStmt(full_attrs, "tirx_hint", IntImm::Int32(1), AsStmt(stmts), source_span), - source_span); -} - } // namespace tirx } // namespace ir_builder } // namespace script diff --git a/src/tirx/script/ir_builder/ir.cc b/src/tirx/script/ir_builder/ir.cc index 675bf816d1d2..b720a71c5f34 100644 --- a/src/tirx/script/ir_builder/ir.cc +++ b/src/tirx/script/ir_builder/ir.cc @@ -513,13 +513,6 @@ ElseFrame Else() { return ElseFrame(n); } -HintFrame Hint(ffi::String message, ffi::Map attrs) { - ffi::ObjectPtr n = ffi::make_object(); - n->message = message; - n->attrs = attrs; - return HintFrame(n); -} - Var EnvThread(ffi::String thread_tag, PrimType dtype) { IterVar iter_var(Range{nullptr}, tvm::PrimVar("", dtype), tvm::tirx::IterVarType::kThreadIndex, thread_tag); @@ -782,7 +775,6 @@ TVM_FFI_STATIC_INIT_BLOCK() { } }) .def("script.ir_builder.tirx.EnvThread", EnvThread) - .def("script.ir_builder.tirx.Hint", Hint) .def("script.ir_builder.tirx.BufferStore", BufferStore) .def("script.ir_builder.tirx.Evaluate", Evaluate) .def("script.ir_builder.tirx.Ptr", Ptr); diff --git a/src/tirx/script/printer/stmt.cc b/src/tirx/script/printer/stmt.cc index 1c2a14b81e92..070cdca0c495 100644 --- a/src/tirx/script/printer/stmt.cc +++ b/src/tirx/script/printer/stmt.cc @@ -786,23 +786,6 @@ TVM_FFI_STATIC_INIT_BLOCK() { rhs = DocsifyLaunchThread(stmt, stmt_p, &define_var, d); } } - if (stmt->attr_key == "tirx_hint") { - if (auto map_node = stmt->node.as>()) { - ffi::Array args; - ffi::Array kwargs_keys; - ffi::Array kwargs_values; - for (const auto& [k, v] : map_node.value()) { - if (k == "message") { - auto s = v.as().value(); - args.push_back(LiteralDoc::Str(s, stmt_p->Attr("node"))); - } else { - kwargs_keys.push_back(k); - kwargs_values.push_back(d->AsDoc(v, stmt_p->Attr("node"))); - } - } - rhs = TIR(d, "hint")->Call(args, kwargs_keys, kwargs_values); - } - } if (!rhs.has_value()) { // Try to collapse consecutive dict-attr-pattern AttrStmts into T.attr({...}) if (IsDictAttrPattern(stmt)) { diff --git a/src/tirx/transform/remove_no_op.cc b/src/tirx/transform/remove_no_op.cc index 99e5c1d928e1..b01d9894916d 100644 --- a/src/tirx/transform/remove_no_op.cc +++ b/src/tirx/transform/remove_no_op.cc @@ -92,9 +92,7 @@ class NoOpRemover : public IRMutatorWithAnalyzer { private: UnchangedOr Mutate_(const AttrStmtNode* op, InplaceMode inplace_mode) final { - if (op->attr_key == "pragma_debug_skip_region") { - return MakeEvaluate(IntImm::Int32(0)); - } else if (op->attr_key == tvm::tirx::attr::async_wait_queue_scope) { + if (op->attr_key == tvm::tirx::attr::async_wait_queue_scope) { auto wait_attrs = GetAsyncWaitAttributes(op); auto wait_cnt = wait_attrs.second; sym::Analyzer ana; diff --git a/src/tirx/transform/storage_rewrite.cc b/src/tirx/transform/storage_rewrite.cc index d6a6d7bdd275..453261fa40c6 100644 --- a/src/tirx/transform/storage_rewrite.cc +++ b/src/tirx/transform/storage_rewrite.cc @@ -240,8 +240,6 @@ class LinearAccessPatternFinder final : public StmtExprVisitor { in_thread_env_ = true; TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(VisitNewScope(op)); in_thread_env_ = false; - } else if (op->attr_key == attr::extern_scope) { - TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(VisitNewScope(op)); } else if (op->attr_key == tvm::tirx::attr::virtual_thread) { TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(VisitNewScope(op)); } else { @@ -391,15 +389,6 @@ class InplaceOpVerifier : public StmtExprVisitor { return std::nullopt; } - ffi::Optional Visit_(const AttrStmtNode* op) final { - // always reject extern code - if (op->attr_key == attr::extern_scope) { - result_ = false; - return std::nullopt; - } - return StmtExprVisitor::Visit_(op); - } - ffi::Optional Visit_(const AllocBufferNode* op) final { // reject inplace for volatile buffers if (op->annotations.count(attr::kVolatile)) { @@ -1091,8 +1080,6 @@ class StoragePlanRewriter : public StmtExprMutator { if (op->attr_key == attr::thread_extent || op->attr_key == tvm::tirx::attr::virtual_thread || attr::IsPragmaKey(op->attr_key)) { PlanNewScope(op); - } else { - TVM_FFI_ICHECK(op->attr_key == attr::extern_scope); } } else if (s.stmt->IsInstance()) { const auto* op = static_cast(s.stmt); diff --git a/tests/python/tirx/test_hint.py b/tests/python/tirx/test_hint.py index a02377f655df..25840ef0e8e1 100644 --- a/tests/python/tirx/test_hint.py +++ b/tests/python/tirx/test_hint.py @@ -14,147 +14,19 @@ # KIND, either express or implied. See the License for the # specific language governing permissions and limitations # under the License. -"""Tests for T.hint() — universal directive primitive for TIRx sketch language.""" - -import tvm_ffi +"""Tests for hint configuration on tile primitive calls.""" import tvm import tvm.script import tvm.testing -from tvm.ir import TensorRegion, assert_structural_equal +from tvm.ir import assert_structural_equal from tvm.script import tirx as T -from tvm.tirx import AttrStmt def from_source(code): return tvm.script.from_source(code, extra_vars={"I": tvm.script.ir, "T": tvm.script.tirx}) -def test_hint_statement(): - """T.hint("msg") as a bare statement produces an AttrStmt with attr_key=tirx_hint.""" - - @T.prim_func - def func(_A: T.Buffer((64,), "float32", scope="global")) -> None: - bx, by, bz = T.cta_id([1, 1, 1]) - warp_id = T.warp_id([1]) - lane_id = T.lane_id([32]) - T.hint("persistent tile scheduler with L2 swizzle") - T.evaluate(0) - - # Walk the IR to find the AttrStmt with tirx_hint - found = [False] - - def visit(stmt): - if isinstance(stmt, AttrStmt) and stmt.attr_key == "tirx_hint": - # node is now a Map with "message" key - assert isinstance(stmt.node, tvm.ir.Map) - assert str(stmt.node["message"]) == "persistent tile scheduler with L2 swizzle" - found[0] = True - - tvm_ffi.structural_walk(func.body, visit) - assert found[0], "Expected AttrStmt with attr_key='tirx_hint' not found" - - -def test_hint_context_manager(): - """with T.hint("msg"): scopes its body inside the AttrStmt.""" - - @T.prim_func - def func(_A: T.Buffer((64,), "float32", scope="global")) -> None: - bx, by, bz = T.cta_id([1, 1, 1]) - warp_id = T.warp_id([1]) - lane_id = T.lane_id([32]) - with T.hint("software pipeline, depth 4"): - T.evaluate(0) - - found = [False] - - def visit(stmt): - if isinstance(stmt, AttrStmt) and stmt.attr_key == "tirx_hint": - assert isinstance(stmt.node, tvm.ir.Map) - assert str(stmt.node["message"]) == "software pipeline, depth 4" - found[0] = True - - tvm_ffi.structural_walk(func.body, visit) - assert found[0], "Expected AttrStmt with attr_key='tirx_hint' not found" - - -def test_hint_with_attrs(): - """T.hint("msg", key="value") passes structured attrs in Map node.""" - - @T.prim_func - def func(_A: T.Buffer((64,), "float32", scope="global")) -> None: - bx, by, bz = T.cta_id([1, 1, 1]) - warp_id = T.warp_id([1]) - lane_id = T.lane_id([32]) - T.hint("scheduler", mode="persistent", depth="4") - T.evaluate(0) - - found = [False] - - def visit(stmt): - if isinstance(stmt, AttrStmt) and stmt.attr_key == "tirx_hint": - assert isinstance(stmt.node, tvm.ir.Map) - assert str(stmt.node["message"]) == "scheduler" - assert str(stmt.node["mode"]) == "persistent" - assert str(stmt.node["depth"]) == "4" - found[0] = True - - tvm_ffi.structural_walk(func.body, visit) - assert found[0], "Expected AttrStmt with attr_key='tirx_hint' not found" - - -def test_hint_printer_roundtrip_statement(): - """Verify T.hint("msg") prints as T.hint("msg") and roundtrips through script/parse.""" - - @T.prim_func - def func(_A: T.Buffer((64,), "float32", scope="global")) -> None: - bx, by, bz = T.cta_id([1, 1, 1]) - warp_id = T.warp_id([1]) - lane_id = T.lane_id([32]) - T.hint("persistent tile scheduler with L2 swizzle") - T.evaluate(0) - - code = func.script() - assert 'hint("persistent tile scheduler with L2 swizzle")' in code - reparsed = from_source(code) - assert_structural_equal(func, reparsed) - - -def test_hint_printer_roundtrip_context_manager(): - """Verify with T.hint("msg"): prints correctly and roundtrips.""" - - @T.prim_func - def func(_A: T.Buffer((64,), "float32", scope="global")) -> None: - bx, by, bz = T.cta_id([1, 1, 1]) - warp_id = T.warp_id([1]) - lane_id = T.lane_id([32]) - with T.hint("software pipeline, depth 4"): - T.evaluate(0) - - code = func.script() - assert 'hint("software pipeline, depth 4")' in code - reparsed = from_source(code) - assert_structural_equal(func, reparsed) - - -def test_hint_printer_roundtrip_with_attrs(): - """Verify T.hint("msg", key="val") prints with kwargs and roundtrips.""" - - @T.prim_func - def func(_A: T.Buffer((64,), "float32", scope="global")) -> None: - bx, by, bz = T.cta_id([1, 1, 1]) - warp_id = T.warp_id([1]) - lane_id = T.lane_id([32]) - T.hint("scheduler", mode="persistent") - T.evaluate(0) - - code = func.script() - assert 'hint("scheduler"' in code - assert 'mode="persistent"' in code - reparsed = from_source(code) - assert_structural_equal(func, reparsed) - - def test_hint_keyword_arg_on_tx_op(): """Tx.op(..., hint="msg") stores hint in TilePrimitiveCall.config.""" from tvm.tirx.buffer import decl_buffer @@ -191,70 +63,5 @@ def func( assert_structural_equal(func, reparsed) -def test_hint_no_message(): - """T.hint(access=...) with no message string.""" - - @T.prim_func - def func(A: T.Buffer((128,), "float32", scope="global")) -> None: - bx, by, bz = T.cta_id([1, 1, 1]) - warp_id = T.warp_id([1]) - lane_id = T.lane_id([32]) - T.hint(access=A[0:64]) - T.evaluate(0) - - found = [False] - - def visit(stmt): - if isinstance(stmt, AttrStmt) and stmt.attr_key == "tirx_hint": - assert isinstance(stmt.node, tvm.ir.Map) - # Should have "access" key but no "message" key - assert "access" in stmt.node - assert "message" not in stmt.node - - assert isinstance(stmt.node["access"], TensorRegion) - found[0] = True - - tvm_ffi.structural_walk(func.body, visit) - assert found[0], "Expected AttrStmt with attr_key='tirx_hint' containing access not found" - - -def test_hint_access_buffer_region(): - """T.hint(access=A[region]) stores the BufferRegion structurally in the IR.""" - - @T.prim_func - def func(A: T.Buffer((128, 64), "float32", scope="global")) -> None: - bx, by, bz = T.cta_id([2, 1, 1]) - warp_id = T.warp_id([1]) - lane_id = T.lane_id([32]) - T.hint("partition", access=A[bx * 64 : (bx + 1) * 64, 0:64]) - T.evaluate(0) - - found = [False] - - def visit(stmt): - if isinstance(stmt, AttrStmt) and stmt.attr_key == "tirx_hint": - assert isinstance(stmt.node, tvm.ir.Map) - assert str(stmt.node["message"]) == "partition" - assert "access" in stmt.node - - assert isinstance(stmt.node["access"], TensorRegion) - br = stmt.node["access"] - assert br.source.name == "A" - assert len(br.region) == 2 - found[0] = True - - tvm_ffi.structural_walk(func.body, visit) - assert found[0], "Expected AttrStmt with structured BufferRegion access not found" - - if __name__ == "__main__": - test_hint_statement() - test_hint_context_manager() - test_hint_with_attrs() - test_hint_printer_roundtrip_statement() - test_hint_printer_roundtrip_context_manager() - test_hint_printer_roundtrip_with_attrs() - test_hint_keyword_arg_on_tx_op() - test_hint_keyword_arg_on_tx_op_roundtrip() - test_hint_no_message() - test_hint_access_buffer_region() + tvm.testing.main()