diff --git a/include/tvm/tir/expr.h b/include/tvm/tir/expr.h index 689b1c0a17ad..3286be7b7959 100644 --- a/include/tvm/tir/expr.h +++ b/include/tvm/tir/expr.h @@ -730,66 +730,6 @@ class ProducerLoad : public PrimExpr { TVM_DEFINE_OBJECT_REF_COW_METHOD(ProducerLoadNode); }; -/*! - * \brief Load the value from buffer_var. - * - * Equivalent to ((DType*)buffer_var)[index] - * where DType is the type specified by type().element_of(). - * - * For example, if type = float32x3, then the load will corresponds to - * - * \code - * - * auto buffer = static_cast(buffer_var); - * auto loaded_val = float32x3(buffer[index.v0], buffer[index.v1], buffer[index.v2]); - * - * \endcode - */ -class LoadNode : public PrimExprNode { - public: - /*! \brief The buffer variable. */ - Var buffer_var; - /*! \brief The index locations to be loaded. */ - PrimExpr index; - /*! \brief The predicate to mask which lanes would be loaded. */ - PrimExpr predicate; - - void VisitAttrs(AttrVisitor* v) { - v->Visit("dtype", &dtype); - v->Visit("buffer_var", &buffer_var); - v->Visit("index", &index); - v->Visit("predicate", &predicate); - v->Visit("span", &span); - } - - bool SEqualReduce(const LoadNode* other, SEqualReducer equal) const { - return equal(dtype, other->dtype) && equal(buffer_var, other->buffer_var) && - equal(index, other->index) && equal(predicate, other->predicate); - } - - void SHashReduce(SHashReducer hash_reduce) const { - hash_reduce(dtype); - hash_reduce(buffer_var); - hash_reduce(index); - hash_reduce(predicate); - } - - static constexpr const char* _type_key = "tir.Load"; - TVM_DECLARE_FINAL_OBJECT_INFO(LoadNode, PrimExprNode); -}; - -/*! - * \brief Managed reference to LoadNode - * \sa LoadNode - */ -class Load : public PrimExpr { - public: - TVM_DLL Load(DataType dtype, Var buffer_var, PrimExpr index, PrimExpr predicate, - Span span = Span()); - TVM_DEFINE_OBJECT_REF_METHODS(Load, PrimExpr, LoadNode); - TVM_DEFINE_OBJECT_REF_COW_METHOD(LoadNode); -}; - /*! * \brief Construct a vector with lanes elements * where its i-th element equals base + i * stride. diff --git a/include/tvm/tir/expr_functor.h b/include/tvm/tir/expr_functor.h index e148d5834f95..3f66164b42c0 100644 --- a/include/tvm/tir/expr_functor.h +++ b/include/tvm/tir/expr_functor.h @@ -120,7 +120,6 @@ class ExprFunctor { } virtual R VisitExpr_(const BufferLoadNode* op, Args... args) EXPR_FUNCTOR_DEFAULT; virtual R VisitExpr_(const ProducerLoadNode* op, Args... args) EXPR_FUNCTOR_DEFAULT; - virtual R VisitExpr_(const LoadNode* op, Args... args) EXPR_FUNCTOR_DEFAULT; virtual R VisitExpr_(const LetNode* op, Args... args) EXPR_FUNCTOR_DEFAULT; virtual R VisitExpr_(const CallNode* op, Args... args) EXPR_FUNCTOR_DEFAULT; virtual R VisitExpr_(const AddNode* op, Args... args) EXPR_FUNCTOR_DEFAULT; @@ -162,7 +161,6 @@ class ExprFunctor { // Set dispatch IR_EXPR_FUNCTOR_DISPATCH(VarNode); IR_EXPR_FUNCTOR_DISPATCH(SizeVarNode); - IR_EXPR_FUNCTOR_DISPATCH(LoadNode); IR_EXPR_FUNCTOR_DISPATCH(BufferLoadNode); IR_EXPR_FUNCTOR_DISPATCH(ProducerLoadNode); IR_EXPR_FUNCTOR_DISPATCH(LetNode); @@ -214,7 +212,6 @@ class TVM_DLL ExprVisitor : public ExprFunctor { // list of functions to override. void VisitExpr_(const VarNode* op) override; void VisitExpr_(const SizeVarNode* op) override; - void VisitExpr_(const LoadNode* op) override; void VisitExpr_(const BufferLoadNode* op) override; void VisitExpr_(const ProducerLoadNode* op) override; void VisitExpr_(const LetNode* op) override; @@ -261,7 +258,6 @@ class TVM_DLL ExprMutator : protected ExprFunctor { // list of functions to override. PrimExpr VisitExpr_(const VarNode* op) override; PrimExpr VisitExpr_(const SizeVarNode* op) override; - PrimExpr VisitExpr_(const LoadNode* op) override; PrimExpr VisitExpr_(const BufferLoadNode* op) override; PrimExpr VisitExpr_(const ProducerLoadNode* op) override; PrimExpr VisitExpr_(const LetNode* op) override; diff --git a/include/tvm/tir/stmt.h b/include/tvm/tir/stmt.h index d7074a7805be..9ed9973871d9 100644 --- a/include/tvm/tir/stmt.h +++ b/include/tvm/tir/stmt.h @@ -213,72 +213,6 @@ class AssertStmt : public Stmt { TVM_DEFINE_OBJECT_REF_COW_METHOD(AssertStmtNode); }; -/*! - * \brief Store value to the buffer. - * - * Equivalent to ((DType*)buffer_var)[index] = value. - * where DType is the type specified by type().element_of(). - * - * For example, if type = float32x3, then the store will corresponds to - * - * \code - * - * auto buffer = static_cast(buffer_var); - * buffer[index.v0] = value.v0; - * buffer[index.v1] = value.v1; - * buffer[index.v2] = value.v2; - * - * \endcode - * \sa LoadNode - */ -class StoreNode : public StmtNode { - public: - /*! \brief The buffer variable. */ - Var buffer_var; - /*! \brief The value to be stored. */ - PrimExpr value; - /*! \brief The index locations to be stored. */ - PrimExpr index; - /*! \brief The predicate to mask which lanes would be stored. */ - PrimExpr predicate; - - void VisitAttrs(AttrVisitor* v) { - v->Visit("buffer_var", &buffer_var); - v->Visit("value", &value); - v->Visit("index", &index); - v->Visit("predicate", &predicate); - v->Visit("span", &span); - } - - bool SEqualReduce(const StoreNode* other, SEqualReducer equal) const { - return equal(buffer_var, other->buffer_var) && equal(value, other->value) && - equal(index, other->index) && equal(predicate, other->predicate); - } - - void SHashReduce(SHashReducer hash_reduce) const { - hash_reduce(buffer_var); - hash_reduce(value); - hash_reduce(index); - hash_reduce(predicate); - } - - static constexpr const char* _type_key = "tir.Store"; - TVM_DECLARE_FINAL_OBJECT_INFO(StoreNode, StmtNode); -}; - -/*! - * \brief Managed reference to StoreNode. - * \sa StoreNode - */ -class Store : public Stmt { - public: - TVM_DLL Store(Var buffer_var, PrimExpr value, PrimExpr index, PrimExpr predicate, - Span span = Span()); - - TVM_DEFINE_OBJECT_REF_METHODS(Store, Stmt, StoreNode); - TVM_DEFINE_OBJECT_REF_COW_METHOD(StoreNode); -}; - /*! * \brief Store value to the high dimension buffer. * diff --git a/include/tvm/tir/stmt_functor.h b/include/tvm/tir/stmt_functor.h index 3adb186fd561..64a384f248e2 100644 --- a/include/tvm/tir/stmt_functor.h +++ b/include/tvm/tir/stmt_functor.h @@ -90,7 +90,6 @@ class StmtFunctor { virtual R VisitStmt_(const AllocateNode* op, Args... args) STMT_FUNCTOR_DEFAULT; virtual R VisitStmt_(const AllocateConstNode* op, Args... args) STMT_FUNCTOR_DEFAULT; virtual R VisitStmt_(const DeclBufferNode* op, Args... args) STMT_FUNCTOR_DEFAULT; - virtual R VisitStmt_(const StoreNode* op, Args... args) STMT_FUNCTOR_DEFAULT; virtual R VisitStmt_(const BufferStoreNode* op, Args... args) STMT_FUNCTOR_DEFAULT; virtual R VisitStmt_(const BufferRealizeNode* op, Args... args) STMT_FUNCTOR_DEFAULT; virtual R VisitStmt_(const AssertStmtNode* op, Args... args) STMT_FUNCTOR_DEFAULT; @@ -117,7 +116,6 @@ class StmtFunctor { IR_STMT_FUNCTOR_DISPATCH(AllocateNode); IR_STMT_FUNCTOR_DISPATCH(AllocateConstNode); IR_STMT_FUNCTOR_DISPATCH(DeclBufferNode); - IR_STMT_FUNCTOR_DISPATCH(StoreNode); IR_STMT_FUNCTOR_DISPATCH(AssertStmtNode); IR_STMT_FUNCTOR_DISPATCH(ProducerStoreNode); IR_STMT_FUNCTOR_DISPATCH(ProducerRealizeNode); @@ -161,7 +159,6 @@ class TVM_DLL StmtVisitor : protected StmtFunctor { void VisitStmt_(const AllocateNode* op) override; void VisitStmt_(const AllocateConstNode* op) override; void VisitStmt_(const DeclBufferNode* op) override; - void VisitStmt_(const StoreNode* op) override; void VisitStmt_(const BufferStoreNode* op) override; void VisitStmt_(const BufferRealizeNode* op) override; void VisitStmt_(const AssertStmtNode* op) override; @@ -263,7 +260,6 @@ class TVM_DLL StmtMutator : protected StmtFunctor { Stmt VisitStmt_(const AllocateNode* op) override; Stmt VisitStmt_(const AllocateConstNode* op) override; Stmt VisitStmt_(const DeclBufferNode* op) override; - Stmt VisitStmt_(const StoreNode* op) override; Stmt VisitStmt_(const BufferStoreNode* op) override; Stmt VisitStmt_(const BufferRealizeNode* op) override; Stmt VisitStmt_(const AssertStmtNode* op) override; diff --git a/python/tvm/ir/json_compact.py b/python/tvm/ir/json_compact.py index a6bcc28dad43..8e9d3550ca4c 100644 --- a/python/tvm/ir/json_compact.py +++ b/python/tvm/ir/json_compact.py @@ -211,7 +211,6 @@ def _convert(item, nodes): "Or": _rename("tir.Or"), "Not": _rename("tir.Not"), "Select": _rename("tir.Select"), - "Load": _rename("tir.Load"), "BufferLoad": _rename("tir.BufferLoad"), "Ramp": _rename("tir.Ramp"), "Broadcast": _rename("tir.Broadcast"), @@ -221,7 +220,6 @@ def _convert(item, nodes): "Any": _rename("tir.Any"), "LetStmt": _rename("tir.LetStmt"), "AssertStmt": _rename("tir.AssertStmt"), - "Store": _rename("tir.Store"), "BufferStore": _rename("tir.BufferStore"), "BufferRealize": _rename("tir.BufferRealize"), "Allocate": _rename("tir.Allocate"), diff --git a/python/tvm/relay/backend/contrib/ethosu/tir/passes.py b/python/tvm/relay/backend/contrib/ethosu/tir/passes.py index f313ff720500..c721efb4710a 100644 --- a/python/tvm/relay/backend/contrib/ethosu/tir/passes.py +++ b/python/tvm/relay/backend/contrib/ethosu/tir/passes.py @@ -49,7 +49,7 @@ def _remove_zero_store(stmt): def _ftransform(f, mod, ctx): return f.with_body( - tvm.tir.stmt_functor.ir_transform(f.body, _remove_zero_store, None, ["tir.Store"]) + tvm.tir.stmt_functor.ir_transform(f.body, _remove_zero_store, None, ["tir.BufferStore"]) ) return tvm.tir.transform.prim_func_pass( diff --git a/python/tvm/relay/backend/contrib/ethosu/tir_to_cs_translator.py b/python/tvm/relay/backend/contrib/ethosu/tir_to_cs_translator.py index ba2c6e209b72..50268f5f874f 100644 --- a/python/tvm/relay/backend/contrib/ethosu/tir_to_cs_translator.py +++ b/python/tvm/relay/backend/contrib/ethosu/tir_to_cs_translator.py @@ -403,7 +403,6 @@ def assign_addresses(buffer_info, npu_ops, scratch_region_map): The key is the buffer name to BufferInfo npu_ops : list A list of Vela NpuOps with tir.BufferLoads for addresses - A list of Vela NpuOps with tir.Loads for addresses scratch_region_map : Dict[tvm.tir.Var, RegionOffset] A buffer_var to region and offset map. Returns diff --git a/python/tvm/script/ir_builder/tir/ir.py b/python/tvm/script/ir_builder/tir/ir.py index 45350c5a65c7..c3ced1e0338b 100644 --- a/python/tvm/script/ir_builder/tir/ir.py +++ b/python/tvm/script/ir_builder/tir/ir.py @@ -63,7 +63,6 @@ FloorMod, IntImm, IterVar, - Load, Max, Min, Mod, @@ -2124,7 +2123,6 @@ def wrapped(*args, **kwargs): "Select", "BufferLoad", "ProducerLoad", - "Load", "Ramp", "Broadcast", "Shuffle", diff --git a/python/tvm/tir/__init__.py b/python/tvm/tir/__init__.py index a77f11862def..10e75b915129 100644 --- a/python/tvm/tir/__init__.py +++ b/python/tvm/tir/__init__.py @@ -24,14 +24,13 @@ from .expr import Var, SizeVar, Reduce, FloatImm, IntImm, StringImm, Cast from .expr import Add, Sub, Mul, Div, Mod, FloorDiv, FloorMod from .expr import Min, Max, EQ, NE, LT, LE, GT, GE, And, Or, Not -from .expr import Select, BufferLoad, ProducerLoad, Load, Ramp, Broadcast, Shuffle +from .expr import Select, BufferLoad, ProducerLoad, Ramp, Broadcast, Shuffle from .expr import Call, CallEffectKind, Let, IterVar, CommReducer, Any from .stmt import Stmt, LetStmt, AssertStmt, ForKind, For, While from .stmt import ( BufferStore, BufferRealize, - Store, ProducerStore, Allocate, AllocateConst, diff --git a/python/tvm/tir/analysis/analysis.py b/python/tvm/tir/analysis/analysis.py index 45b1f745c3de..5feb630e4892 100644 --- a/python/tvm/tir/analysis/analysis.py +++ b/python/tvm/tir/analysis/analysis.py @@ -219,7 +219,8 @@ def calculate_allocated_bytes(func: PrimFunc) -> Dict[str, int]: def detect_buffer_access_lca(func: PrimFunc) -> Dict[Buffer, Stmt]: """Detect the lowest common ancestor(LCA) of buffer access, including both high-level - access(BufferLoad, BufferStore) and low-level access(Load, Store and opaque access). + access (BufferLoad, BufferStore) and low-level access (BufferLoad, BufferStore and opaque + access). The LCA may be a For loop or a Block. Parameters diff --git a/python/tvm/tir/expr.py b/python/tvm/tir/expr.py index cb4a892ac289..52153fd41d63 100644 --- a/python/tvm/tir/expr.py +++ b/python/tvm/tir/expr.py @@ -1013,36 +1013,6 @@ def __init__(self, condition, true_value, false_value, span=None): ) -@tvm._ffi.register_object("tir.Load") -class Load(PrimExprWithOp): - """Load node. - - Parameters - ---------- - dtype : str - The data type. - - buffer_var : Var - The buffer variable in the load expression. - - index : PrimExpr - The index in the load. - - predicate : PrimExpr - The load predicate. - - span : Optional[Span] - The location of this itervar in the source code. - """ - - def __init__(self, dtype, buffer_var, index, predicate=None, span=None): - if predicate is None: - predicate = _ffi_api.const_true(dtype, span) # type: ignore - self.__init_handle_by_constructor__( - _ffi_api.Load, dtype, buffer_var, index, predicate, span # type: ignore - ) - - @tvm._ffi.register_object("tir.BufferLoad") class BufferLoad(PrimExprWithOp): """Buffer load node. diff --git a/python/tvm/tir/stmt.py b/python/tvm/tir/stmt.py index d6cd06a1d915..26b92a46d0dd 100644 --- a/python/tvm/tir/stmt.py +++ b/python/tvm/tir/stmt.py @@ -21,10 +21,10 @@ .. code-block:: python x = tvm.tir.Var("n", "int32") - a = tvm.tir.Var("array", "handle") - st = tvm.tir.stmt.Store(a, x + 1, 1) - assert isinstance(st, tvm.tir.stmt.Store) - assert(st.buffer_var == a) + buffer = tvm.tir.decl_buffer((16,), "float32") + st = tvm.tir.stmt.BufferStore(buffer, 1, (x,)) + assert isinstance(st, tvm.tir.stmt.BufferStore) + assert(st.buffer == buffer) """ from enum import IntEnum from typing import List, Mapping, Optional, Union @@ -189,36 +189,6 @@ def __init__(self, condition, body, span=None): ) -@tvm._ffi.register_object("tir.Store") -class Store(Stmt): - """Store node. - - Parameters - ---------- - buffer_var : Var - The buffer Variable. - - value : PrimExpr - The value we want to store. - - index : PrimExpr - The index in the store expression. - - predicate : PrimExpr - The store predicate. - - span : Optional[Span] - The location of this itervar in the source code. - """ - - def __init__(self, buffer_var, value, index, predicate=None, span=None): - if predicate is None: - predicate = _ffi_api.const_true(value.dtype, span) # type: ignore - self.__init_handle_by_constructor__( - _ffi_api.Store, buffer_var, value, index, predicate, span # type: ignore - ) - - @tvm._ffi.register_object("tir.BufferStore") class BufferStore(Stmt): """Buffer store node. diff --git a/src/contrib/hybrid/codegen_hybrid.cc b/src/contrib/hybrid/codegen_hybrid.cc index 687da61fa019..bde64887856d 100644 --- a/src/contrib/hybrid/codegen_hybrid.cc +++ b/src/contrib/hybrid/codegen_hybrid.cc @@ -250,12 +250,6 @@ void CodeGenHybrid::VisitExpr_(const CallNode* op, std::ostream& os) { // NOLIN } } -void CodeGenHybrid::VisitExpr_(const LoadNode* op, std::ostream& os) { // NOLINT(*) - LOG(FATAL) << "Phase 0 has no Load(s)!"; -} - -void CodeGenHybrid::VisitStmt_(const StoreNode* op) { LOG(FATAL) << "Phase 0 has no Store(s)!"; } - void CodeGenHybrid::VisitExpr_(const BufferLoadNode* op, std::ostream& os) { // NOLINT(*) LOG(FATAL) << "Phase 0 has no BufferLoad(s)!"; } diff --git a/src/contrib/hybrid/codegen_hybrid.h b/src/contrib/hybrid/codegen_hybrid.h index 53026c7fc3b3..d1f578efddd9 100644 --- a/src/contrib/hybrid/codegen_hybrid.h +++ b/src/contrib/hybrid/codegen_hybrid.h @@ -89,7 +89,6 @@ class CodeGenHybrid : public ExprFunctor, } // expression void VisitExpr_(const VarNode* op, std::ostream& os) override; // NOLINT(*) - void VisitExpr_(const LoadNode* op, std::ostream& os) override; // NOLINT(*) void VisitExpr_(const BufferLoadNode* op, std::ostream& os) override; // NOLINT(*) void VisitExpr_(const LetNode* op, std::ostream& os) override; // NOLINT(*) void VisitExpr_(const CallNode* op, std::ostream& os) override; // NOLINT(*) @@ -121,7 +120,6 @@ class CodeGenHybrid : public ExprFunctor, void VisitExpr_(const StringImmNode* op, std::ostream& os) override; // NOLINT(*) // statment void VisitStmt_(const LetStmtNode* op) override; - void VisitStmt_(const StoreNode* op) override; void VisitStmt_(const BufferStoreNode* op) override; void VisitStmt_(const ProducerStoreNode* op) override; void VisitStmt_(const ForNode* op) override; diff --git a/src/relay/printer/text_printer.h b/src/relay/printer/text_printer.h index 707bbec5ad33..a6684bf4e5ce 100644 --- a/src/relay/printer/text_printer.h +++ b/src/relay/printer/text_printer.h @@ -313,7 +313,6 @@ class TIRTextPrinter : public StmtFunctor, Doc VisitExpr_(const SelectNode* op) override; Doc VisitExpr_(const BufferLoadNode* op) override; Doc VisitExpr_(const ProducerLoadNode* op) override; - Doc VisitExpr_(const LoadNode* op) override; Doc VisitExpr_(const RampNode* op) override; Doc VisitExpr_(const BroadcastNode* op) override; Doc VisitExpr_(const tir::LetNode* op) override; @@ -325,7 +324,6 @@ class TIRTextPrinter : public StmtFunctor, Doc VisitStmt_(const LetStmtNode* op) override; Doc VisitStmt_(const AttrStmtNode* op) override; Doc VisitStmt_(const AssertStmtNode* op) override; - Doc VisitStmt_(const StoreNode* op) override; Doc VisitStmt_(const BufferStoreNode* op) override; Doc VisitStmt_(const ProducerStoreNode* op) override; Doc VisitStmt_(const BufferRealizeNode* op) override; diff --git a/src/relay/printer/tir_text_printer.cc b/src/relay/printer/tir_text_printer.cc index eb089bd0d7ed..e9a9ee231358 100644 --- a/src/relay/printer/tir_text_printer.cc +++ b/src/relay/printer/tir_text_printer.cc @@ -379,16 +379,6 @@ Doc TIRTextPrinter::VisitExpr_(const ProducerLoadNode* op) { return doc; } -Doc TIRTextPrinter::VisitExpr_(const LoadNode* op) { - Doc doc; - doc << "(" << PrintDType(op->dtype) << "*)" << Print(op->buffer_var) << "[" << Print(op->index) - << "]"; - if (!is_one(op->predicate)) { - doc << " if " << Print(op->predicate); - } - return doc; -} - Doc TIRTextPrinter::VisitExpr_(const RampNode* op) { Doc doc; doc << "ramp(" << Print(op->base) << ", " << Print(op->stride) << ", " << op->lanes << ")"; @@ -477,15 +467,6 @@ Doc TIRTextPrinter::VisitStmt_(const AssertStmtNode* op) { return doc; } -Doc TIRTextPrinter::VisitStmt_(const StoreNode* op) { - Doc doc; - doc << Print(op->buffer_var) << "[" << Print(op->index) << "] = " << Print(op->value); - if (!is_one(op->predicate)) { - doc << " if " << Print(op->predicate); - } - return doc; -} - Doc TIRTextPrinter::VisitStmt_(const BufferStoreNode* op) { Doc doc; doc << Print(op->buffer) << Print(op->indices) << " = " << Print(op->value); diff --git a/src/relay/printer/tvmscript_printer.cc b/src/relay/printer/tvmscript_printer.cc index 096611095097..b0085b82426e 100644 --- a/src/relay/printer/tvmscript_printer.cc +++ b/src/relay/printer/tvmscript_printer.cc @@ -242,7 +242,6 @@ class TVMScriptPrinter : public StmtFunctor, Doc VisitExpr_(const StringImmNode* op, ExprPrecedence* out_precedence) override; Doc VisitExpr_(const ProducerLoadNode* op, ExprPrecedence* out_precedence) override; Doc VisitExpr_(const BufferLoadNode* op, ExprPrecedence* out_precedence) override; - Doc VisitExpr_(const LoadNode* op, ExprPrecedence* out_precedence) override; Doc VisitExpr_(const RampNode* op, ExprPrecedence* out_precedence) override; Doc VisitExpr_(const BroadcastNode* op, ExprPrecedence* out_precedence) override; Doc VisitExpr_(const tir::LetNode* op, ExprPrecedence* out_precedence) override; @@ -254,7 +253,6 @@ class TVMScriptPrinter : public StmtFunctor, Doc VisitStmt_(const LetStmtNode* op) override; Doc VisitStmt_(const AttrStmtNode* op) override; Doc VisitStmt_(const AssertStmtNode* op) override; - Doc VisitStmt_(const StoreNode* op) override; Doc VisitStmt_(const BufferStoreNode* op) override; Doc VisitStmt_(const BufferRealizeNode* op) override; Doc VisitStmt_(const AllocateNode* op) override; @@ -910,23 +908,6 @@ Doc TVMScriptPrinter::VisitExpr_(const BufferLoadNode* op, ExprPrecedence* out_p return doc; } -Doc TVMScriptPrinter::VisitExpr_(const LoadNode* op, ExprPrecedence* out_precedence) { - *out_precedence = ExprPrecedence::kIdentity; - Doc doc; - if (op->dtype == DataType::Float(32) && is_one(op->predicate) && - op->buffer_var->dtype == DataType::Float(32)) { - doc << Print(op->buffer_var) << "[" << Print(op->index) << "]"; - } else { - doc << tir_prefix_ << ".load(" << PrintDType(op->dtype) << ", " << Print(op->buffer_var) << ", " - << Print(op->index); - if (!is_one(op->predicate) || op->dtype.lanes() != 1) { - doc << ", " << Print(op->predicate); - } - doc << ")"; - } - return doc; -} - Doc TVMScriptPrinter::VisitExpr_(const RampNode* op, ExprPrecedence* out_precedence) { *out_precedence = ExprPrecedence::kIdentity; Doc doc; @@ -1078,13 +1059,6 @@ Doc TVMScriptPrinter::VisitStmt_(const AssertStmtNode* op) { return doc; } -Doc TVMScriptPrinter::VisitStmt_(const StoreNode* op) { - Doc doc; - doc << tir_prefix_ << ".store(" << Print(op->buffer_var) << ", " << Print(op->index) << ", " - << Print(op->value) << ", " << Print(op->predicate) << ")"; - return doc; -} - Doc TVMScriptPrinter::VisitStmt_(const BufferRealizeNode* op) { LOG(FATAL) << "TVM Script Printer Internal Error: All the BufferRealize should be folded with Attr"; diff --git a/src/script/printer/legacy_repr.cc b/src/script/printer/legacy_repr.cc index 2909e059f3e3..01fb514c497e 100644 --- a/src/script/printer/legacy_repr.cc +++ b/src/script/printer/legacy_repr.cc @@ -459,18 +459,6 @@ TVM_STATIC_IR_FUNCTOR(ReprLegacyPrinter, vtable) (*p) << ")"; }); -TVM_STATIC_IR_FUNCTOR(ReprLegacyPrinter, vtable) - .set_dispatch([](const ObjectRef& node, ReprLegacyPrinter* p) { - auto* op = static_cast(node.get()); - (*p) << op->buffer_var << "["; - p->Print(op->index); - (*p) << "]"; - if (!is_one(op->predicate)) { - (*p) << " if "; - p->Print(op->predicate); - } - }); - TVM_STATIC_IR_FUNCTOR(ReprLegacyPrinter, vtable) .set_dispatch([](const ObjectRef& node, ReprLegacyPrinter* p) { auto* op = static_cast(node.get()); @@ -664,21 +652,6 @@ TVM_STATIC_IR_FUNCTOR(ReprLegacyPrinter, vtable) (*p) << "}\n"; }); -TVM_STATIC_IR_FUNCTOR(ReprLegacyPrinter, vtable) - .set_dispatch([](const ObjectRef& node, ReprLegacyPrinter* p) { - auto* op = static_cast(node.get()); - p->PrintIndent(); - (*p) << op->buffer_var << "["; - p->Print(op->index); - (*p) << "] = "; - p->Print(op->value); - if (!is_one(op->predicate)) { - (*p) << " if "; - p->Print(op->predicate); - } - (*p) << '\n'; - }); - TVM_STATIC_IR_FUNCTOR(ReprLegacyPrinter, vtable) .set_dispatch([](const ObjectRef& node, ReprLegacyPrinter* p) { auto* op = static_cast(node.get()); diff --git a/src/script/printer/tir/expr.cc b/src/script/printer/tir/expr.cc index 9c4f62eb1c1b..710f2eab22e2 100644 --- a/src/script/printer/tir/expr.cc +++ b/src/script/printer/tir/expr.cc @@ -288,11 +288,6 @@ TVM_STATIC_IR_FUNCTOR(IRDocsifier, vtable) LOG(FATAL) << "ValueError: Reduce should never exist in TIR: " << r; }); -TVM_STATIC_IR_FUNCTOR(IRDocsifier, vtable) - .set_dispatch("", [](tir::Load load, ObjectPath p, IRDocsifier d) -> Doc { - LOG(FATAL) << "ValueError: Load has been deprecated for BufferLoad: " << load; - }); - #define TVM_SCRIPT_PRINTER_DEF_BINARY(NodeType, OpString) \ TVM_STATIC_IR_FUNCTOR(IRDocsifier, vtable) \ .set_dispatch("", \ @@ -393,7 +388,6 @@ TVM_SCRIPT_REPR(tir::CommReducerNode, ReprPrintTIR); TVM_SCRIPT_REPR(tir::IndexMapNode, ReprPrintTIR); TVM_SCRIPT_REPR(tir::AnyNode, ReprPrintTIR); TVM_SCRIPT_REPR(tir::ReduceNode, ReprPrintTIR); -TVM_SCRIPT_REPR(tir::LoadNode, ReprPrintTIR); } // namespace printer } // namespace script diff --git a/src/script/printer/tir/stmt.cc b/src/script/printer/tir/stmt.cc index 591d1e3bc1da..384ad6a9407b 100644 --- a/src/script/printer/tir/stmt.cc +++ b/src/script/printer/tir/stmt.cc @@ -419,12 +419,6 @@ TVM_STATIC_IR_FUNCTOR(IRDocsifier, vtable) return DoConciseScoping(lhs, rhs.value(), &(*f)->stmts, concise); }); -TVM_STATIC_IR_FUNCTOR(IRDocsifier, vtable) - .set_dispatch( // - "", [](tir::Store stmt, ObjectPath p, IRDocsifier d) -> Doc { - LOG(FATAL) << "ValueError: Store has been deprecated for BufferStore: " << stmt; - }); - TVM_SCRIPT_REPR(tir::LetStmtNode, ReprPrintTIR); TVM_SCRIPT_REPR(tir::AttrStmtNode, ReprPrintTIR); TVM_SCRIPT_REPR(tir::AssertStmtNode, ReprPrintTIR); @@ -437,7 +431,6 @@ TVM_SCRIPT_REPR(tir::SeqStmtNode, ReprPrintTIR); TVM_SCRIPT_REPR(tir::IfThenElseNode, ReprPrintTIR); TVM_SCRIPT_REPR(tir::EvaluateNode, ReprPrintTIR); TVM_SCRIPT_REPR(tir::BufferRealizeNode, ReprPrintTIR); -TVM_SCRIPT_REPR(tir::StoreNode, ReprPrintTIR); } // namespace printer } // namespace script diff --git a/src/target/llvm/codegen_llvm.cc b/src/target/llvm/codegen_llvm.cc index dd9c3ddb5b2f..365beb5f5dd0 100644 --- a/src/target/llvm/codegen_llvm.cc +++ b/src/target/llvm/codegen_llvm.cc @@ -1566,10 +1566,6 @@ llvm::Value* CodeGenLLVM::VisitExpr_(const LetNode* op) { return MakeValue(op->body); } -llvm::Value* CodeGenLLVM::VisitExpr_(const LoadNode* op) { - LOG(FATAL) << "Unexpected deprecated LoadNode. Use BufferLoadNode instead."; -} - bool CodeGenLLVM::HasAlignmentPadding(DataType dtype) { const llvm::DataLayout& data_layout = module_->getDataLayout(); int bytes = data_layout.getTypeAllocSize(DTypeToLLVMType(dtype)); @@ -1766,10 +1762,6 @@ llvm::Value* CodeGenLLVM::VisitExpr_(const BroadcastNode* op) { return CreateBroadcast(MakeValue(op->value), op->lanes); } -void CodeGenLLVM::VisitStmt_(const StoreNode* op) { - LOG(FATAL) << "Unexpected deprecated StoreNode. Use BufferStoreNode instead."; -} - void CodeGenLLVM::VisitStmt_(const BufferStoreNode* op) { EmitDebugLocation(op); DataType value_dtype = op->value.dtype(); diff --git a/src/target/llvm/codegen_llvm.h b/src/target/llvm/codegen_llvm.h index 62b0b0cc4bd8..b46ae07b8403 100644 --- a/src/target/llvm/codegen_llvm.h +++ b/src/target/llvm/codegen_llvm.h @@ -204,14 +204,12 @@ class CodeGenLLVM : public ExprFunctor, llvm::Value* VisitExpr_(const NotNode* op) override; llvm::Value* VisitExpr_(const SelectNode* op) override; llvm::Value* VisitExpr_(const LetNode* op) override; - llvm::Value* VisitExpr_(const LoadNode* op) override; llvm::Value* VisitExpr_(const BufferLoadNode* op) override; llvm::Value* VisitExpr_(const CallNode* op) override; llvm::Value* VisitExpr_(const RampNode* op) override; llvm::Value* VisitExpr_(const ShuffleNode* op) override; llvm::Value* VisitExpr_(const BroadcastNode* op) override; // stmt - void VisitStmt_(const StoreNode* op) override; void VisitStmt_(const BufferStoreNode* op) override; void VisitStmt_(const ForNode* op) override; void VisitStmt_(const WhileNode* op) override; diff --git a/src/target/source/codegen_c.cc b/src/target/source/codegen_c.cc index a4332476f335..7807b46e4227 100644 --- a/src/target/source/codegen_c.cc +++ b/src/target/source/codegen_c.cc @@ -666,10 +666,6 @@ void CodeGenC::VisitStmt_(const AllocateConstNode* op) { void CodeGenC::VisitStmt_(const DeclBufferNode* op) { this->PrintStmt(op->body); } -void CodeGenC::VisitExpr_(const LoadNode* op, std::ostream& os) { // NOLINT(*) - LOG(FATAL) << "Unexpected deprecated LoadNode. Use BufferLoadNode instead."; -} - void CodeGenC::VisitExpr_(const BufferLoadNode* op, std::ostream& os) { // NOLINT(*) ICHECK_EQ(op->indices.size(), 1) << "Load from non-flat memory not supported."; @@ -729,10 +725,6 @@ void CodeGenC::VisitExpr_(const BufferLoadNode* op, std::ostream& os) { // NOLI } } -void CodeGenC::VisitStmt_(const StoreNode* op) { - LOG(FATAL) << "Unexpected deprecated StoreNode. Use BufferStoreNode instead."; -} - void CodeGenC::VisitStmt_(const BufferStoreNode* op) { ICHECK_EQ(op->indices.size(), 1) << "Store to non-flat memory not supported."; diff --git a/src/target/source/codegen_c.h b/src/target/source/codegen_c.h index 40733808d61b..4f0da5a9dbad 100644 --- a/src/target/source/codegen_c.h +++ b/src/target/source/codegen_c.h @@ -126,7 +126,6 @@ class CodeGenC : public ExprFunctor, virtual void InitFuncState(const PrimFunc& f); // expression void VisitExpr_(const VarNode* op, std::ostream& os) override; // NOLINT(*) - void VisitExpr_(const LoadNode* op, std::ostream& os) override; // NOLINT(*) void VisitExpr_(const BufferLoadNode* op, std::ostream& os) override; // NOLINT(*) void VisitExpr_(const LetNode* op, std::ostream& os) override; // NOLINT(*) void VisitExpr_(const CallNode* op, std::ostream& os) override; // NOLINT(*) @@ -156,7 +155,6 @@ class CodeGenC : public ExprFunctor, void VisitExpr_(const StringImmNode* op, std::ostream& os) override; // NOLINT(*) // statment void VisitStmt_(const LetStmtNode* op) override; - void VisitStmt_(const StoreNode* op) override; void VisitStmt_(const BufferStoreNode* op) override; void VisitStmt_(const ForNode* op) override; void VisitStmt_(const WhileNode* op) override; diff --git a/src/target/source/codegen_opencl.cc b/src/target/source/codegen_opencl.cc index 89cc09aeadb2..613b1d084701 100644 --- a/src/target/source/codegen_opencl.cc +++ b/src/target/source/codegen_opencl.cc @@ -382,10 +382,6 @@ std::string CodeGenOpenCL::CastTo(std::string value, DataType target) { return os.str(); } -void CodeGenOpenCL::VisitStmt_(const StoreNode* op) { - LOG(FATAL) << "Unexpected use of deprecated StoreNode. Please use BufferStoreNode instead."; -} - void CodeGenOpenCL::VisitStmt_(const BufferStoreNode* op) { if (auto call = op->value.as()) { if (call->op.same_as(builtin::texture2d_load())) { diff --git a/src/target/source/codegen_opencl.h b/src/target/source/codegen_opencl.h index 169759976119..05734b6a54eb 100644 --- a/src/target/source/codegen_opencl.h +++ b/src/target/source/codegen_opencl.h @@ -68,7 +68,6 @@ class CodeGenOpenCL final : public CodeGenC { void VisitExpr_(const CallNode* op, std::ostream& os) final; // NOLINT(*) void VisitExpr_(const CastNode* op, std::ostream& os) final; // NOLINT(*) void VisitExpr_(const FloatImmNode* op, std::ostream& os) final; // NOLINT(*) - void VisitStmt_(const StoreNode* op) final; // NOLINT(*) void VisitStmt_(const BufferStoreNode* op) final; // NOLINT(*) // overload min and max to avoid ambiguous call errors diff --git a/src/target/stackvm/codegen_stackvm.cc b/src/target/stackvm/codegen_stackvm.cc index eac9ad849419..15e483562a4b 100644 --- a/src/target/stackvm/codegen_stackvm.cc +++ b/src/target/stackvm/codegen_stackvm.cc @@ -139,10 +139,6 @@ int CodeGenStackVM::GetVarID(const VarNode* v) const { return it->second; } -void CodeGenStackVM::VisitExpr_(const LoadNode* op) { - LOG(FATAL) << "Unexpected use of deprecated LoadNode. Please use BufferLoadNode instead."; -} - void CodeGenStackVM::VisitExpr_(const BufferLoadNode* op) { ICHECK_EQ(op->indices.size(), 1) << "StackVM expects flat 1-d buffers. " << "Has StorageFlatten (TE-based schedules) or " @@ -162,10 +158,6 @@ void CodeGenStackVM::VisitExpr_(const BufferLoadNode* op) { } } -void CodeGenStackVM::VisitStmt_(const StoreNode* op) { - LOG(FATAL) << "Unexpected use of deprecated StoreNode. Please use BufferStoreNode instead."; -} - void CodeGenStackVM::VisitStmt_(const BufferStoreNode* op) { ICHECK_EQ(op->indices.size(), 1) << "StackVM expects flat 1-d buffers. " << "Has StorageFlatten (TE-based schedules) or " diff --git a/src/target/stackvm/codegen_stackvm.h b/src/target/stackvm/codegen_stackvm.h index ae6f316b475d..87a61f76cc61 100644 --- a/src/target/stackvm/codegen_stackvm.h +++ b/src/target/stackvm/codegen_stackvm.h @@ -107,7 +107,6 @@ class CodeGenStackVM : public ExprFunctor, // overloadable functions // expression void VisitExpr_(const VarNode* op) final; - void VisitExpr_(const LoadNode* op) final; void VisitExpr_(const BufferLoadNode* op) final; void VisitExpr_(const LetNode* op) final; void VisitExpr_(const CallNode* op) final; @@ -136,7 +135,6 @@ class CodeGenStackVM : public ExprFunctor, void VisitExpr_(const StringImmNode* op) final; // statment void VisitStmt_(const LetStmtNode* op) final; - void VisitStmt_(const StoreNode* op) final; void VisitStmt_(const BufferStoreNode* op) final; void VisitStmt_(const ForNode* op) final; void VisitStmt_(const IfThenElseNode* op) final; diff --git a/src/te/autodiff/jacobian.cc b/src/te/autodiff/jacobian.cc index a77688b43efb..78788dbe1a0c 100644 --- a/src/te/autodiff/jacobian.cc +++ b/src/te/autodiff/jacobian.cc @@ -75,7 +75,6 @@ class JacobianMutator : public ExprMutator { } } - PrimExpr VisitExpr_(const LoadNode* op) NOT_IMPLEMENTED; PrimExpr VisitExpr_(const LetNode* op) NOT_IMPLEMENTED; PrimExpr VisitExpr_(const ProducerLoadNode* op) final { diff --git a/src/tir/analysis/block_access_region_detector.cc b/src/tir/analysis/block_access_region_detector.cc index ab328efaa6d1..409356c2b155 100644 --- a/src/tir/analysis/block_access_region_detector.cc +++ b/src/tir/analysis/block_access_region_detector.cc @@ -107,9 +107,7 @@ class BlockReadWriteDetector : public StmtExprVisitor { void VisitStmt_(const IfThenElseNode* op) override; void VisitStmt_(const BlockRealizeNode* op) override; void VisitStmt_(const BufferStoreNode* op) override; - void VisitStmt_(const StoreNode* op) override; void VisitExpr_(const BufferLoadNode* op) override; - void VisitExpr_(const LoadNode* op) override; void VisitExpr_(const VarNode* op) override; void VisitExpr_(const CallNode* op) override; }; @@ -144,10 +142,6 @@ Array BlockReadWriteDetector::CollectOpaques() { void BlockReadWriteDetector::VisitExpr_(const VarNode* op) { UpdateOpaque(GetRef(op)); } -void BlockReadWriteDetector::VisitExpr_(const LoadNode* op) { - LOG(FATAL) << "Unexpected use of deprecated LoadNode. Please use BufferLoadNode instead."; -} - void BlockReadWriteDetector::VisitExpr_(const BufferLoadNode* op) { std::vector relaxed_region; for (const PrimExpr& index : op->indices) { @@ -224,10 +218,6 @@ void BlockReadWriteDetector::VisitExpr_(const CallNode* op) { StmtExprVisitor::VisitExpr_(op); } -void BlockReadWriteDetector::VisitStmt_(const StoreNode* op) { - LOG(FATAL) << "Unexpected use of deprecated StoreNode. Please use BufferStoreNode instead."; -} - void BlockReadWriteDetector::VisitStmt_(const BufferStoreNode* op) { std::vector relaxed_region; for (const PrimExpr& index : op->indices) { diff --git a/src/tir/analysis/buffer_access_lca_detector.cc b/src/tir/analysis/buffer_access_lca_detector.cc index 64d10fae2ff1..ff0b11a73c9b 100644 --- a/src/tir/analysis/buffer_access_lca_detector.cc +++ b/src/tir/analysis/buffer_access_lca_detector.cc @@ -237,15 +237,6 @@ class LCADetector : public StmtExprVisitor { // Works for Load/Store and opaque access. void VisitExpr_(const VarNode* op) final { VisitBufferVar(op); } - // Explict to visit buffer data in Load and Store node. - void VisitExpr_(const LoadNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated LoadNode. Please use BufferLoadNode instead."; - } - - void VisitStmt_(const StoreNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated StoreNode. Please use BufferStoreNode instead."; - } - void VisitBufferVar(const VarNode* op) { auto it = buffer_var_map_.find(op); if (it != buffer_var_map_.end()) { diff --git a/src/tir/analysis/device_constraint_utils.cc b/src/tir/analysis/device_constraint_utils.cc index d0933e0691dd..4554038bc770 100644 --- a/src/tir/analysis/device_constraint_utils.cc +++ b/src/tir/analysis/device_constraint_utils.cc @@ -249,15 +249,6 @@ class ApplyDeviceConstraintsMutator : public StmtExprMutator { private: PrimExpr VisitExpr_(const VarNode* var_node) final { return Subst(var_node); } - PrimExpr VisitExpr_(const LoadNode* load_node) final { - Load new_load = Downcast(StmtExprMutator::VisitExpr_(load_node)); - Var new_buffer_var = Subst(new_load->buffer_var.get()); - if (!new_buffer_var.same_as(new_load->buffer_var)) { - return Load(load_node->dtype, new_buffer_var, load_node->index, load_node->predicate); - } - return std::move(new_load); - } - PrimExpr VisitExpr_(const BufferLoadNode* buffer_load_node) final { BufferLoad new_buffer_load = Downcast(StmtExprMutator::VisitExpr_(buffer_load_node)); @@ -296,15 +287,6 @@ class ApplyDeviceConstraintsMutator : public StmtExprMutator { return StmtExprMutator::VisitStmt_(allocate_node); } - Stmt VisitStmt_(const StoreNode* store_node) final { - Store new_store = Downcast(StmtExprMutator::VisitStmt_(store_node)); - Var new_buffer_var = Subst(new_store->buffer_var.get()); - if (!new_buffer_var.same_as(new_store->buffer_var)) { - Store(new_buffer_var, new_store->value, new_store->index, new_store->predicate); - } - return std::move(new_store); - } - Stmt VisitStmt_(const BufferStoreNode* buffer_store_node) final { BufferStore new_buffer_store = Downcast(StmtExprMutator::VisitStmt_(buffer_store_node)); diff --git a/src/tir/analysis/side_effect.cc b/src/tir/analysis/side_effect.cc index 5613961e2b66..7c5d39283774 100644 --- a/src/tir/analysis/side_effect.cc +++ b/src/tir/analysis/side_effect.cc @@ -37,11 +37,6 @@ class ExprSideEffect : public ExprVisitor { ExprVisitor::VisitExpr(e); } - void VisitExpr_(const LoadNode* op) final { - this->UpdateEffect(CallEffectKind::kReadState); - ExprVisitor::VisitExpr_(op); - } - void VisitExpr_(const BufferLoadNode* op) final { this->UpdateEffect(CallEffectKind::kReadState); ExprVisitor::VisitExpr_(op); diff --git a/src/tir/analysis/var_touch.cc b/src/tir/analysis/var_touch.cc index f92afc4d15a1..8c2ed6c43255 100644 --- a/src/tir/analysis/var_touch.cc +++ b/src/tir/analysis/var_touch.cc @@ -44,14 +44,6 @@ class VarTouchVisitor : public StmtExprVisitor { void VisitExpr_(const VarNode* op) final { Handle(op); } - void VisitExpr_(const LoadNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated LoadNode. Please use BufferLoadNode instead."; - } - - void VisitStmt_(const StoreNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated StoreNode. Please use BufferStoreNode instead."; - } - void VisitStmt_(const BufferStoreNode* op) final { Handle(op->buffer->data.get()); StmtVisitor::VisitStmt_(op); diff --git a/src/tir/analysis/var_use_def_analysis.cc b/src/tir/analysis/var_use_def_analysis.cc index 7ef8e532a396..9d1105cd15c7 100644 --- a/src/tir/analysis/var_use_def_analysis.cc +++ b/src/tir/analysis/var_use_def_analysis.cc @@ -72,10 +72,6 @@ void VarUseDefAnalyzer::VisitStmt_(const AllocateConstNode* op) { StmtExprVisitor::VisitStmt_(op); } -void VarUseDefAnalyzer::VisitStmt_(const StoreNode* op) { - LOG(FATAL) << "Unexpected use of deprecated StoreNode. Please use BufferStoreNode instead."; -} - void VarUseDefAnalyzer::VisitStmt_(const BufferStoreNode* op) { VisitBuffer(op->buffer); StmtExprVisitor::VisitStmt_(op); @@ -112,10 +108,6 @@ void VarUseDefAnalyzer::VisitExpr_(const ReduceNode* op) { StmtExprVisitor::VisitExpr_(op); } -void VarUseDefAnalyzer::VisitExpr_(const LoadNode* op) { - LOG(FATAL) << "Unexpected use of deprecated LoadNode. Please use BufferLoadNode instead."; -} - void VarUseDefAnalyzer::VisitExpr_(const BufferLoadNode* op) { VisitBuffer(op->buffer); StmtExprVisitor::VisitExpr_(op); diff --git a/src/tir/analysis/var_use_def_analysis.h b/src/tir/analysis/var_use_def_analysis.h index ad275011d90c..5d0ceed13c5e 100644 --- a/src/tir/analysis/var_use_def_analysis.h +++ b/src/tir/analysis/var_use_def_analysis.h @@ -62,8 +62,6 @@ class VarUseDefAnalyzer : public StmtExprVisitor { void VisitStmt_(const AllocateConstNode* op) final; - void VisitStmt_(const StoreNode* op) final; - void VisitStmt_(const BufferStoreNode* op) final; void VisitExpr_(const LetNode* op) final; @@ -72,8 +70,6 @@ class VarUseDefAnalyzer : public StmtExprVisitor { void VisitExpr_(const ReduceNode* op) final; - void VisitExpr_(const LoadNode* op) final; - void VisitExpr_(const BufferLoadNode* op) final; void HandleDef(const VarNode* v); diff --git a/src/tir/analysis/verify_gpu_code.cc b/src/tir/analysis/verify_gpu_code.cc index 3377515a9589..3d6c66d0e193 100644 --- a/src/tir/analysis/verify_gpu_code.cc +++ b/src/tir/analysis/verify_gpu_code.cc @@ -186,14 +186,6 @@ class GPUCodeVerifier : public StmtExprVisitor { StmtVisitor::VisitStmt_(op); } - void VisitExpr_(const LoadNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated LoadNode. Please use BufferLoadNode instead."; - } - - void VisitStmt_(const StoreNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated StoreNode. Please use BufferStoreNode instead."; - } - void CheckBufferIndicesVectorizable(const Array indices) { for (const auto index : indices) { if (const auto* ramp = index.as()) { diff --git a/src/tir/analysis/verify_memory.cc b/src/tir/analysis/verify_memory.cc index 9d932d236355..a210a555b4cd 100644 --- a/src/tir/analysis/verify_memory.cc +++ b/src/tir/analysis/verify_memory.cc @@ -88,14 +88,6 @@ class MemoryAccessVerifier final : protected StmtExprVisitor { } } - void VisitExpr_(const LoadNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated LoadNode. Please use BufferLoadNode instead."; - } - - void VisitStmt_(const StoreNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated StoreNode. Please use BufferStoreNode instead."; - } - void VisitExpr_(const BufferLoadNode* op) final { HandleLoadStoreToVariable(op->buffer->data); return StmtExprVisitor::VisitExpr_(op); diff --git a/src/tir/ir/expr.cc b/src/tir/ir/expr.cc index db09ac17e6eb..f1b69f38e42e 100644 --- a/src/tir/ir/expr.cc +++ b/src/tir/ir/expr.cc @@ -422,73 +422,6 @@ TVM_REGISTER_GLOBAL("tir.Select") TVM_REGISTER_NODE_TYPE(SelectNode); -// Load -Load::Load(DataType dtype, Var buffer_var, PrimExpr index, PrimExpr predicate, Span span) { - LOG(FATAL) << "Unexpected use of deprecated Store node for buffer " << buffer_var->name_hint - << ". Use BufferStore instead."; - ICHECK(buffer_var.defined()); - ICHECK(predicate.defined()); - ICHECK(index.defined()); - - // Assume that the array elements have 1 lane, unless a type - // annotation tells us otherwise. - int element_lanes = 1; - auto pointer_type = tir::GetPointerType(buffer_var->type_annotation); - if (pointer_type.has_value()) { - // Cannot check element type of array, as it may be different than - // the loaded type in some cases. - // - // 1. Booleans use DataType::Int(8) while stored, and the codegens - // handle cast to boolean. - // - // 2. The StorageRewrite pass can merge multiple allocations at - // the same scope, regardless of element type. The codegen is - // then responsible for casting to the output type. - - // TODO(Lunderberg): Uncomment this check once it can be applied. - // See https://discuss.tvm.apache.org/t/pre-rfc-vectorized-tir-buffers/10615 - // for discussion. - - // ICHECK(dtype.element_of() == pointer_type->element_of()) - // << "Type mismatch, cannot load type " << dtype << " from buffer " << - // buffer_var->name_hint - // << " of type " << pointer_type.value(); - element_lanes = pointer_type->lanes(); - } - - // The C-based codegens assume that all loads occur on a array with - // non-vectorized elements, and cast between - // vectorized/non-vectorized arrays as needed. Ideally, these - // should be changed to explicit casts in the TIR graph, rather than - // being handled at the code-gen level. - ICHECK((dtype.lanes() == element_lanes * index.dtype().lanes()) || - (dtype.lanes() == index.dtype().lanes())); - ICHECK((dtype.lanes() == element_lanes * predicate.dtype().lanes()) || - (dtype.lanes() == index.dtype().lanes())); - - ObjectPtr node = make_object(); - node->dtype = dtype; - node->buffer_var = std::move(buffer_var); - node->index = std::move(index); - node->predicate = std::move(predicate); - node->span = std::move(span); - - data_ = std::move(node); -} - -TVM_REGISTER_GLOBAL("tir.Load").set_body([](TVMArgs args, TVMRetValue* ret) { - DataType t = args[0]; - if (args.size() == 3) { - *ret = Load(t, args[1], args[2], const_true(t.lanes()), Span()); - } else if (args.size() == 4) { - *ret = Load(t, args[1], args[2], args[3], Span()); - } else { - *ret = Load(t, args[1], args[2], args[3], args[4]); - } -}); - -TVM_REGISTER_NODE_TYPE(LoadNode); - // Ramp Ramp::Ramp(PrimExpr base, PrimExpr stride, int lanes, Span span) { ICHECK(base.defined()); diff --git a/src/tir/ir/expr_functor.cc b/src/tir/ir/expr_functor.cc index b3b09e54f2e2..8a93d9dd8242 100644 --- a/src/tir/ir/expr_functor.cc +++ b/src/tir/ir/expr_functor.cc @@ -34,10 +34,6 @@ void ExprVisitor::VisitExpr_(const SizeVarNode* op) { void ExprVisitor::VisitExpr_(const AnyNode* op) {} -void ExprVisitor::VisitExpr_(const LoadNode* op) { - LOG(FATAL) << "Unexpected use of deprecated LoadNode. Please use BufferLoadNode instead."; -} - void ExprVisitor::VisitExpr_(const BufferLoadNode* op) { VisitArray(op->indices, [this](const PrimExpr& e) { this->VisitExpr(e); }); } @@ -125,10 +121,6 @@ PrimExpr ExprMutator::VisitExpr_(const SizeVarNode* op) { PrimExpr ExprMutator::VisitExpr_(const AnyNode* op) { return GetRef(op); } -PrimExpr ExprMutator::VisitExpr_(const LoadNode* op) { - LOG(FATAL) << "Unexpected use of deprecated LoadNode. Please use BufferLoadNode instead."; -} - PrimExpr ExprMutator::VisitExpr_(const BufferLoadNode* op) { auto fmutate = [this](const PrimExpr& e) { return this->VisitExpr(e); }; Array indices = op->indices.Map(fmutate); diff --git a/src/tir/ir/stmt.cc b/src/tir/ir/stmt.cc index fd2a98554da6..e4569898f7ed 100644 --- a/src/tir/ir/stmt.cc +++ b/src/tir/ir/stmt.cc @@ -195,59 +195,6 @@ TVM_REGISTER_GLOBAL("tir.While").set_body_typed([](PrimExpr condition, Stmt body TVM_REGISTER_NODE_TYPE(WhileNode); -// Store -Store::Store(Var buffer_var, PrimExpr value, PrimExpr index, PrimExpr predicate, Span span) { - LOG(FATAL) << "Unexpected use of deprecated Store node for buffer " << buffer_var->name_hint - << ". Use BufferStore instead."; - ICHECK(value.defined()); - ICHECK(index.defined()); - ICHECK(predicate.defined()); - - // Assume that the array elements have 1 lane, unless a type - // annotation tells us otherwise. - int element_lanes = 1; - auto pointer_type = tir::GetPointerType(buffer_var->type_annotation); - if (pointer_type.has_value()) { - // Currently cannot check element type of array, see Load::Load - // for details. - - // TODO(Lunderberg): Uncomment this check once it can be applied. - // See https://discuss.tvm.apache.org/t/pre-rfc-vectorized-tir-buffers/10615 - // for discussion. - - // ICHECK_EQ(value.dtype().element_of(), pointer_type->element_of()) - // << "Type mismatch, cannot store type " << value.dtype() << " into buffer " - // << buffer_var->name_hint << " of type " << pointer_type.value(); - element_lanes = pointer_type->lanes(); - } - - ICHECK((value.dtype().lanes() == element_lanes * index.dtype().lanes()) || - (value.dtype().lanes() == index.dtype().lanes())); - ICHECK((value.dtype().lanes() == element_lanes * predicate.dtype().lanes()) || - (value.dtype().lanes() == index.dtype().lanes())); - - ObjectPtr node = make_object(); - node->buffer_var = std::move(buffer_var); - node->value = std::move(value); - node->index = std::move(index); - node->predicate = std::move(predicate); - node->span = std::move(span); - data_ = std::move(node); -} - -TVM_REGISTER_GLOBAL("tir.Store").set_body([](TVMArgs args, TVMRetValue* ret) { - PrimExpr value = args[1]; - if (args.size() == 3) { - *ret = Store(args[0], value, args[2], const_true(value.dtype().lanes()), Span()); - } else if (args.size() == 4) { - *ret = Store(args[0], value, args[2], args[3], Span()); - } else { - *ret = Store(args[0], value, args[2], args[3], args[4]); - } -}); - -TVM_REGISTER_NODE_TYPE(StoreNode); - // ProducerStore ProducerStore::ProducerStore(DataProducer producer, PrimExpr value, Array indices, Span span) { diff --git a/src/tir/ir/stmt_functor.cc b/src/tir/ir/stmt_functor.cc index 1d00b8bd364f..22f8aa06d56c 100644 --- a/src/tir/ir/stmt_functor.cc +++ b/src/tir/ir/stmt_functor.cc @@ -66,10 +66,6 @@ void StmtVisitor::VisitStmt_(const AllocateConstNode* op) { void StmtVisitor::VisitStmt_(const DeclBufferNode* op) { this->VisitStmt(op->body); } -void StmtVisitor::VisitStmt_(const StoreNode* op) { - LOG(FATAL) << "Unexpected use of deprecated StoreNode. Please use BufferStoreNode instead."; -} - void StmtVisitor::VisitStmt_(const BufferStoreNode* op) { this->VisitExpr(op->value); VisitArray(op->indices, [this](const PrimExpr& e) { this->VisitExpr(e); }); @@ -369,10 +365,6 @@ Stmt StmtMutator::VisitStmt_(const IfThenElseNode* op) { } } -Stmt StmtMutator::VisitStmt_(const StoreNode* op) { - LOG(FATAL) << "Unexpected use of deprecated StoreNode. Please use BufferStoreNode instead."; -} - Stmt StmtMutator::VisitStmt_(const BufferStoreNode* op) { PrimExpr value = this->VisitExpr(op->value); Array indices = Internal::Mutate(this, op->indices); @@ -673,14 +665,6 @@ class IRSubstitute : public StmtExprMutator { return std::move(var); } - PrimExpr VisitExpr_(const LoadNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated LoadNode. Please use BufferLoadNode instead."; - } - - Stmt VisitStmt_(const StoreNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated StoreNode. Please use BufferStoreNode instead."; - } - PrimExpr VisitExpr_(const BufferLoadNode* op) final { auto node = Downcast(StmtExprMutator::VisitExpr_(op)); return VisitBufferAccess(std::move(node)); diff --git a/src/tir/schedule/analysis/reducer.cc b/src/tir/schedule/analysis/reducer.cc index 5f1a0e355608..b8ba7ac96f3c 100644 --- a/src/tir/schedule/analysis/reducer.cc +++ b/src/tir/schedule/analysis/reducer.cc @@ -62,24 +62,6 @@ class PatternMatcher : public ExprVisitor { } } - void VisitExpr_(const LoadNode* op) final { - const auto* ptr = expr_to_match_.as(); - if (ptr == nullptr) { - match_success_ = false; - } else { - if (!op->buffer_var.same_as(ptr->buffer_var)) { - match_success_ = false; - } else { - PrimExpr tmp = expr_to_match_; - expr_to_match_ = ptr->predicate; - VisitExpr(op->predicate); - expr_to_match_ = ptr->index; - VisitExpr(op->index); - std::swap(expr_to_match_, tmp); - } - } - } - void VisitExpr_(const LetNode* op) final { const auto* ptr = expr_to_match_.as(); if (ptr == nullptr) { diff --git a/src/tir/schedule/primitive/cache_index.cc b/src/tir/schedule/primitive/cache_index.cc index 0316feefd5de..7e80a2a2e9de 100644 --- a/src/tir/schedule/primitive/cache_index.cc +++ b/src/tir/schedule/primitive/cache_index.cc @@ -425,10 +425,6 @@ class CacheIndexRewriter : public StmtExprMutator { return ret_stmt; } - PrimExpr VisitExpr_(const LoadNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated LoadNode. Please use BufferLoadNode instead."; - } - private: /*! \brief The parent scope of the insertion */ const StmtSRef& scope_sref_; diff --git a/src/tir/schedule/primitive/cache_read_write.cc b/src/tir/schedule/primitive/cache_read_write.cc index 39e915ba961a..bc18e3f8fccf 100644 --- a/src/tir/schedule/primitive/cache_read_write.cc +++ b/src/tir/schedule/primitive/cache_read_write.cc @@ -843,10 +843,6 @@ class CacheReadRewriter : public StmtExprMutator { return ExprMutator::VisitExpr_(load); } - PrimExpr VisitExpr_(const LoadNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated LoadNode. Please use BufferLoadNode instead."; - } - PrimExpr VisitExpr_(const VarNode* op) final { if (op == info_->read_buffer->data.get()) { return info_->write_buffer->data; @@ -1067,14 +1063,6 @@ class CacheWriteRewriter : public StmtExprMutator { return ExprMutator::VisitExpr_(load); } - PrimExpr VisitExpr_(const LoadNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated LoadNode. Please use BufferLoadNode instead."; - } - - Stmt VisitStmt_(const StoreNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated StoreNode. Please use BufferStoreNode instead."; - } - PrimExpr VisitExpr_(const VarNode* op) final { if (op == info_->write_buffer->data.get()) { return info_->read_buffer->data; diff --git a/src/tir/schedule/primitive/compute_inline.cc b/src/tir/schedule/primitive/compute_inline.cc index ad4aa9ef748e..b64351186ac5 100644 --- a/src/tir/schedule/primitive/compute_inline.cc +++ b/src/tir/schedule/primitive/compute_inline.cc @@ -261,14 +261,6 @@ class BaseInliner : public StmtExprMutator { return StmtExprMutator::VisitExpr_(var); } - PrimExpr VisitExpr_(const LoadNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated LoadNode. Please use BufferLoadNode instead."; - } - - Stmt VisitStmt_(const StoreNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated StoreNode. Please use BufferStoreNode instead."; - } - Stmt VisitStmt_(const ForNode* loop) final { if (src_stmt.get() == loop) { loop = tgt_stmt.as(); diff --git a/src/tir/transforms/bf16_legalize.cc b/src/tir/transforms/bf16_legalize.cc index 8c5982e80916..3b89558622e9 100644 --- a/src/tir/transforms/bf16_legalize.cc +++ b/src/tir/transforms/bf16_legalize.cc @@ -258,10 +258,6 @@ class BF16LowerRewriter : public StmtExprMutator { } } - Stmt VisitStmt_(const StoreNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated StoreNode. Please use BufferStoreNode instead."; - } - PrimExpr VisitExpr_(const BufferLoadNode* op) final { PrimExpr ret = StmtExprMutator::VisitExpr_(op); op = ret.as(); @@ -274,10 +270,6 @@ class BF16LowerRewriter : public StmtExprMutator { } } - PrimExpr VisitExpr_(const LoadNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated LoadNode. Please use BufferLoadNode instead."; - } - PrimExpr VisitExpr_(const FloatImmNode* op) final { if (op->dtype.is_bfloat16()) { return IntImm(DataType::UInt(16, op->dtype.lanes()), diff --git a/src/tir/transforms/bound_checker.cc b/src/tir/transforms/bound_checker.cc index 5a4178a018cf..f5aa6773e66d 100644 --- a/src/tir/transforms/bound_checker.cc +++ b/src/tir/transforms/bound_checker.cc @@ -78,14 +78,6 @@ class BoundChecker : public StmtExprMutator { return StmtExprMutator::VisitExpr_(op); } - PrimExpr VisitExpr_(const LoadNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated LoadNode. Please use BufferLoadNode instead."; - } - - Stmt VisitStmt_(const StoreNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated StoreNode. Please use BufferStoreNode instead."; - } - Stmt VisitStmt_(const BufferStoreNode* op) final { store_scope_bound_collector_.clear(); process_store_ = true; diff --git a/src/tir/transforms/common_subexpr_elim.cc b/src/tir/transforms/common_subexpr_elim.cc index acda9220b731..b6c52ec1a3be 100644 --- a/src/tir/transforms/common_subexpr_elim.cc +++ b/src/tir/transforms/common_subexpr_elim.cc @@ -69,8 +69,7 @@ namespace tir { bool CommonSubexpressionEliminator::ForbiddenComputation(const PrimExpr& expr) { // Function calls, loads and buffer loads are absolutely forbidden as introducing them into // variables would change the semantics of the program. - return (expr.as() != nullptr || expr.as() != nullptr || - expr.as() != nullptr); + return (expr.as() != nullptr || expr.as() != nullptr); } /*! @@ -116,7 +115,7 @@ bool CommonSubexpressionEliminator::CanContainEligibleComputations(const PrimExp // not harm the indexing mode of the CPU, but as we are still far from ASM code, we // finally want to perform such simplifications, which tend to happen fairly frequently. - // return ( (expr.as() == nullptr) && (expr.as() == nullptr) ) + // return (expr.as() == nullptr) return true; } diff --git a/src/tir/transforms/compact_buffer_region.cc b/src/tir/transforms/compact_buffer_region.cc index b517150ce9f4..3cfb1a474002 100644 --- a/src/tir/transforms/compact_buffer_region.cc +++ b/src/tir/transforms/compact_buffer_region.cc @@ -134,14 +134,6 @@ class BufferAccessRegionCollector : public StmtExprVisitor { void VisitExpr_(const VarNode* op) final { VisitBufferVar(GetRef(op)); } - void VisitExpr_(const LoadNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated LoadNode. Please use BufferLoadNode instead."; - } - - void VisitStmt_(const StoreNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated StoreNode. Please use BufferStoreNode instead."; - } - void VisitStmt_(const ForNode* op) final { ancestor_loops_.push_back(op); Range loop_range = Range::FromMinExtent(op->min, op->extent); diff --git a/src/tir/transforms/coproc_sync.cc b/src/tir/transforms/coproc_sync.cc index 69913f4bd604..65ee33d2dad6 100644 --- a/src/tir/transforms/coproc_sync.cc +++ b/src/tir/transforms/coproc_sync.cc @@ -38,12 +38,6 @@ namespace tir { // Visitor to find touched set by co-processor scope. class CoProcTouchedBuffer : public StmtExprVisitor { public: - void VisitExpr_(const LoadNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated LoadNode. Please use BufferLoadNode instead."; - } - void VisitStmt_(const StoreNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated StoreNode. Please use BufferStoreNode instead."; - } void VisitExpr_(const BufferLoadNode* op) final { if (in_scope_) { touched_[op->buffer->data.get()].coproc = true; diff --git a/src/tir/transforms/inject_copy_intrin.cc b/src/tir/transforms/inject_copy_intrin.cc index 81842ff808ff..f7b14f49977f 100644 --- a/src/tir/transforms/inject_copy_intrin.cc +++ b/src/tir/transforms/inject_copy_intrin.cc @@ -94,7 +94,7 @@ class CopyIntrinInjector : public StmtMutator { load = cast->value.as(); } if (load == nullptr) { - *error_info = "the 'LoadNode' of body is a nullptr."; + *error_info = "the 'BufferLoadNode' of body is a nullptr."; return false; } if (load->dtype.lanes() != 1) return false; diff --git a/src/tir/transforms/inject_double_buffer.cc b/src/tir/transforms/inject_double_buffer.cc index 91052cbf572d..c99264041efb 100644 --- a/src/tir/transforms/inject_double_buffer.cc +++ b/src/tir/transforms/inject_double_buffer.cc @@ -170,14 +170,6 @@ class DoubleBufferInjector : public StmtExprMutator { return stmt; } - PrimExpr VisitExpr_(const LoadNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated LoadNode. Please use BufferLoadNode instead."; - } - - Stmt VisitStmt_(const StoreNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated StoreNode. Please use BufferStoreNode instead."; - } - Stmt VisitStmt_(const BufferStoreNode* op) final { auto node = Downcast(StmtExprMutator::VisitStmt_(op)); diff --git a/src/tir/transforms/inject_virtual_thread.cc b/src/tir/transforms/inject_virtual_thread.cc index 5b54b8abee8e..1094abd9e1fa 100644 --- a/src/tir/transforms/inject_virtual_thread.cc +++ b/src/tir/transforms/inject_virtual_thread.cc @@ -50,9 +50,6 @@ class ExprTouched final : public StmtExprVisitor { if (expr_touched_ && !check_write_) return; StmtExprVisitor::VisitStmt(n); } - void VisitExpr_(const LoadNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated StoreNode. Please use BufferStoreNode instead."; - } void VisitExpr_(const BufferLoadNode* op) final { HandleUseVar(op->buffer->data.get()); StmtExprVisitor::VisitExpr_(op); @@ -106,10 +103,6 @@ class VarTouchedAnalysis : public StmtVisitor { this->VisitStmt(op->body); } - void VisitStmt_(const StoreNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated StoreNode. Please use BufferStoreNode instead."; - } - void VisitStmt_(const BufferStoreNode* op) final { ExprTouched tc(touched_var_, false); tc(op->value); @@ -244,14 +237,6 @@ class VTInjector : public arith::IRMutatorWithAnalyzer { trigger_base_inject_ = !allow_share_; return StmtExprMutator::VisitStmt_(op); } - // Load - PrimExpr VisitExpr_(const LoadNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated LoadNode. Please use BufferLoadNode instead."; - } - // Store - Stmt VisitStmt_(const StoreNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated StoreNode. Please use BufferStoreNode instead."; - } // BufferLoad PrimExpr VisitExpr_(const BufferLoadNode* op) final { auto node = Downcast(StmtExprMutator::VisitExpr_(op)); diff --git a/src/tir/transforms/install_debug_spans.h b/src/tir/transforms/install_debug_spans.h index c71891aba5a6..40f3e07940cf 100644 --- a/src/tir/transforms/install_debug_spans.h +++ b/src/tir/transforms/install_debug_spans.h @@ -75,7 +75,6 @@ X(Allocate) \ X(AllocateConst) \ X(DeclBuffer) \ - X(Store) \ X(BufferStore) \ X(BufferRealize) \ X(AssertStmt) \ diff --git a/src/tir/transforms/ir_utils.cc b/src/tir/transforms/ir_utils.cc index afd7ba43cf93..f6e4ac45c612 100644 --- a/src/tir/transforms/ir_utils.cc +++ b/src/tir/transforms/ir_utils.cc @@ -110,14 +110,6 @@ class IRConvertSSA final : public StmtExprMutator { } } - PrimExpr VisitExpr_(const LoadNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated LoadNode. Please use BufferLoadNode instead."; - } - - Stmt VisitStmt_(const StoreNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated StoreNode. Please use BufferStoreNode instead."; - } - PrimExpr VisitExpr_(const BufferLoadNode* op) final { auto node = Downcast(StmtExprMutator::VisitExpr_(op)); auto output = VisitBufferAccess(std::move(node)); diff --git a/src/tir/transforms/lower_custom_datatypes.cc b/src/tir/transforms/lower_custom_datatypes.cc index 241b656ace6c..8480189855b8 100644 --- a/src/tir/transforms/lower_custom_datatypes.cc +++ b/src/tir/transforms/lower_custom_datatypes.cc @@ -103,14 +103,6 @@ class CustomDatatypesLowerer : public StmtExprMutator { } } - PrimExpr VisitExpr_(const LoadNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated LoadNode. Please use BufferLoadNode instead."; - } - - Stmt VisitStmt_(const StoreNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated StoreNode. Please use BufferStoreNode instead."; - } - PrimExpr VisitExpr_(const BufferLoadNode* op) final { auto node = Downcast(StmtExprMutator::VisitExpr_(op)); auto modified = VisitBufferAccess(node); diff --git a/src/tir/transforms/lower_match_buffer.cc b/src/tir/transforms/lower_match_buffer.cc index 2aa6d18b4d11..700587fe0e21 100644 --- a/src/tir/transforms/lower_match_buffer.cc +++ b/src/tir/transforms/lower_match_buffer.cc @@ -117,20 +117,6 @@ class MatchBufferLower : public StmtExprMutator { } } - PrimExpr VisitExpr_(const LoadNode* op) final { - PrimExpr expr = StmtExprMutator::VisitExpr_(op); - CHECK(var_map_.find(op->buffer_var) == var_map_.end()) - << "Load from buffer created by match_buffer is not allowed, but got: " << expr; - return expr; - } - - Stmt VisitStmt_(const StoreNode* op) final { - Stmt stmt = StmtExprMutator::VisitStmt_(op); - CHECK(var_map_.find(op->buffer_var) == var_map_.end()) - << "Store from buffer created by match_buffer is not allowed, but got: " << stmt; - return stmt; - } - BufferRegion VisitBufferRegion(const BufferRegion& buffer_region) { const Buffer& buffer = buffer_region->buffer; auto it = match_buffers_.find(buffer); diff --git a/src/tir/transforms/lower_thread_allreduce.cc b/src/tir/transforms/lower_thread_allreduce.cc index cade9a90566d..9104b7d51bf9 100644 --- a/src/tir/transforms/lower_thread_allreduce.cc +++ b/src/tir/transforms/lower_thread_allreduce.cc @@ -109,14 +109,6 @@ class ThreadAllreduceBuilder final : public StmtExprMutator { } } - PrimExpr VisitExpr_(const LoadNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated LoadNode. Please use BufferLoadNode instead."; - } - - Stmt VisitStmt_(const StoreNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated StoreNode. Please use BufferStoreNode instead."; - } - PrimExpr VisitExpr_(const BufferLoadNode* op) final { { auto it = load_remap_.find(op->buffer->data.get()); diff --git a/src/tir/transforms/lower_warp_memory.cc b/src/tir/transforms/lower_warp_memory.cc index 9d2ff88540fc..571f512bfd14 100644 --- a/src/tir/transforms/lower_warp_memory.cc +++ b/src/tir/transforms/lower_warp_memory.cc @@ -125,10 +125,6 @@ class WarpStoreCoeffFinder : private StmtExprVisitor { StmtExprVisitor::VisitExpr_(op); } - void VisitStmt_(const StoreNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated StoreNode. Please use BufferStoreNode instead."; - } - void VisitStmt_(const BufferStoreNode* op) final { if (op->buffer->data.get() != buffer_) { StmtVisitor::VisitStmt_(op); @@ -293,14 +289,6 @@ class WarpAccessRewriter : protected StmtExprMutator { return StmtExprMutator::VisitExpr_(op); } - Stmt VisitStmt_(const StoreNode* op) override { - LOG(FATAL) << "Unexpected use of deprecated StoreNode. Please use BufferStoreNode instead."; - } - - PrimExpr VisitExpr_(const LoadNode* op) override { - LOG(FATAL) << "Unexpected use of deprecated LoadNode. Please use BufferLoadNode instead."; - } - Stmt VisitStmt_(const BufferStoreNode* op) override { auto store = Downcast(StmtExprMutator::VisitStmt_(op)); diff --git a/src/tir/transforms/merge_dynamic_shared_memory_allocations.cc b/src/tir/transforms/merge_dynamic_shared_memory_allocations.cc index eab660e2a47b..02cfad3fca4a 100644 --- a/src/tir/transforms/merge_dynamic_shared_memory_allocations.cc +++ b/src/tir/transforms/merge_dynamic_shared_memory_allocations.cc @@ -103,10 +103,6 @@ class DynSharedMemLinearAccessPatternFinder final : public StmtExprVisitor { StmtExprVisitor::VisitStmt_(op); } - void VisitStmt_(const StoreNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated StoreNode. Please use BufferStoreNode instead."; - } - void VisitStmt_(const BufferStoreNode* op) final { scope_.push_back(StmtEntry()); // visit subexpr @@ -140,10 +136,6 @@ class DynSharedMemLinearAccessPatternFinder final : public StmtExprVisitor { } } - void VisitExpr_(const LoadNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated LoadNode. Please use BufferLoadNode instead."; - } - void VisitExpr_(const BufferLoadNode* op) final { // Add write access. StmtExprVisitor::VisitExpr_(op); @@ -307,14 +299,6 @@ class DynamicSharedMemoryRewriter : public StmtExprMutator { return StmtExprMutator::VisitStmt_(op); } - PrimExpr VisitExpr_(const LoadNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated LoadNode. Please use BufferLoadNode instead."; - } - - Stmt VisitStmt_(const StoreNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated StoreNode. Please use BufferStoreNode instead."; - } - PrimExpr VisitExpr_(const BufferLoadNode* op) final { auto node = Downcast(StmtExprMutator::VisitExpr_(op)); return VisitBufferAccess(std::move(node)); diff --git a/src/tir/transforms/narrow_datatype.cc b/src/tir/transforms/narrow_datatype.cc index ad8132521d47..7b6187af64b8 100644 --- a/src/tir/transforms/narrow_datatype.cc +++ b/src/tir/transforms/narrow_datatype.cc @@ -34,7 +34,7 @@ namespace tvm { namespace tir { -// This pass narrows indexing expressions (like StoreNode::Index) +// This pass narrows indexing expressions (like BufferStoreNode::indices) // that trivially fit into i32/i16 (denoted by `target_bits_`) to // i32/i16. Considering that i32/i16 indices may be more // efficient on some backends (while i64 may be more efficient @@ -223,14 +223,6 @@ class NarrowDataTypeRewriter : public IndexDataTypeRewriter { using Parent::VisitExpr_; using Parent::VisitStmt_; - Stmt VisitStmt_(const StoreNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated StoreNode. Please use BufferStoreNode instead."; - } - - PrimExpr VisitExpr_(const LoadNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated LoadNode. Please use BufferLoadNode instead."; - } - PrimExpr VisitExpr_(const VarNode* op) final { if (auto it = visitor_.vmap.find(op); !var_remap_.count(op) && it != visitor_.vmap.end()) { var_remap_[op] = Var(op->name_hint, it->second); diff --git a/src/tir/transforms/renew_defs.cc b/src/tir/transforms/renew_defs.cc index 90399f7a0586..7eac6645239e 100644 --- a/src/tir/transforms/renew_defs.cc +++ b/src/tir/transforms/renew_defs.cc @@ -157,14 +157,6 @@ class RenewDefMutator : public StmtExprMutator { } } - PrimExpr VisitExpr_(const LoadNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated LoadNode. Please use BufferLoadNode instead."; - } - - Stmt VisitStmt_(const StoreNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated StoreNode. Please use BufferStoreNode instead."; - } - private: Var ReDefineVar(const Var& var) { Var new_var = Var(make_object(*var.get())); diff --git a/src/tir/transforms/rewrite_unsafe_select.cc b/src/tir/transforms/rewrite_unsafe_select.cc index b4082e2040fd..21660991ab37 100644 --- a/src/tir/transforms/rewrite_unsafe_select.cc +++ b/src/tir/transforms/rewrite_unsafe_select.cc @@ -67,9 +67,6 @@ class UnsafeExprDetector : public ExprFunctor { // Load is considered unsafe. return true; } - bool VisitExpr_(const LoadNode* op) { - LOG(FATAL) << "Unexpected use of deprecated LoadNode. Please use BufferLoadNode instead."; - } bool VisitExpr_(const AddNode* op) final { return BinaryOp(op); } bool VisitExpr_(const SubNode* op) final { return BinaryOp(op); } bool VisitExpr_(const MulNode* op) final { return BinaryOp(op); } diff --git a/src/tir/transforms/simplify.cc b/src/tir/transforms/simplify.cc index 7dd52f941c46..cc088e8f74c6 100644 --- a/src/tir/transforms/simplify.cc +++ b/src/tir/transforms/simplify.cc @@ -207,10 +207,6 @@ class StmtSimplifier : public IRMutatorWithAnalyzer { return Parent::VisitExpr_(op); } - Stmt VisitStmt_(const StoreNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated StoreNode. Please use BufferStoreNode instead."; - } - // eliminate useless stores Stmt VisitStmt_(const BufferStoreNode* op) final { BufferStore store = Downcast(Parent::VisitStmt_(op)); diff --git a/src/tir/transforms/storage_access.cc b/src/tir/transforms/storage_access.cc index 8729ab1ed296..b34cfdfb3128 100644 --- a/src/tir/transforms/storage_access.cc +++ b/src/tir/transforms/storage_access.cc @@ -33,14 +33,6 @@ namespace tvm { namespace tir { -void StorageAccessVisitor::VisitExpr_(const LoadNode* op) { - LOG(FATAL) << "Unexpected use of deprecated LoadNode. Please use BufferLoadNode instead."; -} - -void StorageAccessVisitor::VisitStmt_(const StoreNode* op) { - LOG(FATAL) << "Unexpected use of deprecated StoreNode. Please use BufferStoreNode instead."; -} - void StorageAccessVisitor::VisitExpr_(const BufferLoadNode* op) { Var buf = op->buffer->data; StorageScope scope = GetScope(buf); diff --git a/src/tir/transforms/storage_access.h b/src/tir/transforms/storage_access.h index ac64e2f5cb65..8fac0c302a70 100644 --- a/src/tir/transforms/storage_access.h +++ b/src/tir/transforms/storage_access.h @@ -81,8 +81,6 @@ class StorageAccessVisitor : public StmtExprVisitor { std::vector access; }; // override visitor pattern - void VisitExpr_(const LoadNode* op) final; - void VisitStmt_(const StoreNode* op) final; void VisitExpr_(const BufferLoadNode* op) final; void VisitStmt_(const BufferStoreNode* op) final; void VisitStmt_(const EvaluateNode* op) final; diff --git a/src/tir/transforms/storage_flatten.cc b/src/tir/transforms/storage_flatten.cc index 58f4eba83893..9d6e6a35b875 100644 --- a/src/tir/transforms/storage_flatten.cc +++ b/src/tir/transforms/storage_flatten.cc @@ -846,14 +846,6 @@ class BufferBindUnwrapper : public StmtExprMutator { return output; } - Stmt VisitStmt_(const StoreNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated StoreNode. Please use BufferStoreNode instead."; - } - - PrimExpr VisitExpr_(const LoadNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated LoadNode. Please use BufferLoadNode instead."; - } - Stmt VisitStmt_(const AttrStmtNode* op) final { ICHECK_NE(op->attr_key, attr::buffer_dim_align) << "BufferBindUnwrapper assumes that all buffers have accurate strides, " @@ -1386,14 +1378,6 @@ class StorageFlattener : public StmtExprMutator { cache_line_size_ = cache_line_size; } - Stmt VisitStmt_(const StoreNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated StoreNode. Please use BufferStoreNode instead."; - } - - PrimExpr VisitExpr_(const LoadNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated LoadNode. Please use BufferLoadNode instead."; - } - Stmt VisitStmt_(const AttrStmtNode* op) final { ICHECK_NE(op->attr_key, attr::buffer_dim_align) << "StorageFlattener assumes that all buffers have accurate strides, " diff --git a/src/tir/transforms/storage_rewrite.cc b/src/tir/transforms/storage_rewrite.cc index 7e09bda70371..5d08f6e45be1 100644 --- a/src/tir/transforms/storage_rewrite.cc +++ b/src/tir/transforms/storage_rewrite.cc @@ -100,10 +100,6 @@ class LinearAccessPatternFinder final : public StmtExprVisitor { StmtExprVisitor::VisitStmt_(op); } - void VisitStmt_(const StoreNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated StoreNode. Please use BufferStoreNode instead."; - } - void VisitStmt_(const BufferStoreNode* op) final { scope_.push_back(StmtEntry()); // visit subexpr @@ -129,10 +125,6 @@ class LinearAccessPatternFinder final : public StmtExprVisitor { } } - void VisitExpr_(const LoadNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated LoadNode. Please use BufferLoadNode instead."; - } - void VisitExpr_(const BufferLoadNode* op) final { // Add write access. StmtExprVisitor::VisitExpr_(op); @@ -280,8 +272,6 @@ class InplaceOpVerifier : public StmtExprVisitor { VisitStmt_(static_cast(stmt)); } else if (stmt->IsInstance()) { VisitStmt_(static_cast(stmt)); - } else if (stmt->IsInstance()) { - VisitStmt_(static_cast(stmt)); } else if (stmt->IsInstance()) { VisitStmt_(static_cast(stmt)); } else { @@ -309,10 +299,6 @@ class InplaceOpVerifier : public StmtExprVisitor { } } - void VisitStmt_(const StoreNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated StoreNode. Please use BufferStoreNode instead."; - } - void VisitStmt_(const BufferStoreNode* op) final { ++mem_nest_; for (const auto& index : op->indices) { @@ -337,10 +323,6 @@ class InplaceOpVerifier : public StmtExprVisitor { StmtExprVisitor::VisitStmt_(op); } - void VisitExpr_(const LoadNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated LoadNode. Please use BufferLoadNode instead."; - } - void VisitExpr_(const BufferLoadNode* op) final { const VarNode* buf = op->buffer->data.get(); // cannot read from dst_ (no reduction) @@ -416,14 +398,6 @@ class StoragePlanRewriter : public StmtExprMutator { return stmt; } - Stmt VisitStmt_(const StoreNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated StoreNode. Please use BufferStoreNode instead."; - } - - PrimExpr VisitExpr_(const LoadNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated LoadNode. Please use BufferLoadNode instead."; - } - template Node VisitBufferAccess(Node node) { auto it = alloc_map_.find(node->buffer->data.get()); @@ -1149,14 +1123,6 @@ class VectorTypeAccessChecker : public StmtExprVisitor { } } - void VisitExpr_(const LoadNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated LoadNode. Please use BufferLoadNode instead."; - } - - void VisitStmt_(const StoreNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated StoreNode. Please use BufferStoreNode instead."; - } - void VisitExpr_(const BufferLoadNode* op) final { OnArrayAccess(op->dtype, op->buffer->data.get(), op->indices); StmtExprVisitor::VisitExpr_(op); @@ -1414,14 +1380,6 @@ class VectorTypeRewriter : public StmtExprMutator { } } - PrimExpr VisitExpr_(const LoadNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated LoadNode. Please use BufferLoadNode instead."; - } - - Stmt VisitStmt_(const StoreNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated StoreNode. Please use BufferStoreNode instead."; - } - template Node VisitBufferAccess(Node node) { if (!rewrite_indices_) { diff --git a/src/tir/transforms/thread_storage_sync.cc b/src/tir/transforms/thread_storage_sync.cc index 4becd8ffd74f..90d7bedcf97b 100644 --- a/src/tir/transforms/thread_storage_sync.cc +++ b/src/tir/transforms/thread_storage_sync.cc @@ -314,13 +314,6 @@ class ThreadSyncInserter : public StmtExprMutator { return StmtExprMutator::VisitStmt(stmt); } } - PrimExpr VisitExpr_(const LoadNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated LoadNode. Please use BufferLoadNode instead."; - } - - Stmt VisitStmt_(const StoreNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated StoreNode. Please use BufferStoreNode instead."; - } PrimExpr VisitExpr_(const BufferLoadNode* op) final { if (sync_scope_.rank == StorageRank::kGlobal && GetScope(op->buffer->data).rank == StorageRank::kGlobal) { diff --git a/src/tir/transforms/unroll_loop.cc b/src/tir/transforms/unroll_loop.cc index dc14e4512f1e..e43d4d7fd69e 100644 --- a/src/tir/transforms/unroll_loop.cc +++ b/src/tir/transforms/unroll_loop.cc @@ -161,10 +161,6 @@ class LoopUnroller : public StmtExprMutator { } } - Stmt VisitStmt_(const StoreNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated StoreNode. Please use BufferStoreNode instead."; - } - PrimExpr VisitExpr_(const BufferLoadNode* op) final { if (unroll_local_access_) { auto storage_scope = runtime::StorageScope::Create(GetPtrStorageScope(op->buffer->data)); diff --git a/src/tir/transforms/update_pointer_storage_scope.cc b/src/tir/transforms/update_pointer_storage_scope.cc index 3a9e4717241d..157e81d77f12 100644 --- a/src/tir/transforms/update_pointer_storage_scope.cc +++ b/src/tir/transforms/update_pointer_storage_scope.cc @@ -94,19 +94,11 @@ Buffer UpdatePointerStorageScope::GetUpdatedBuffer(Buffer buf) { return buf; } -PrimExpr UpdatePointerStorageScope::VisitExpr_(const LoadNode* op) { - LOG(FATAL) << "Unexpected use of deprecated LoadNode. Please use BufferLoadNode instead."; -} - PrimExpr UpdatePointerStorageScope::VisitExpr_(const BufferLoadNode* op) { auto node = Downcast(StmtExprMutator::VisitExpr_(op)); return UpdateBufferAccess(node); } -Stmt UpdatePointerStorageScope::VisitStmt_(const StoreNode* op) { - LOG(FATAL) << "Unexpected use of deprecated StoreNode. Please use BufferStoreNode instead."; -} - Stmt UpdatePointerStorageScope::VisitStmt_(const BufferStoreNode* op) { auto node = Downcast(StmtExprMutator::VisitStmt_(op)); return UpdateBufferAccess(node); diff --git a/src/tir/transforms/update_pointer_storage_scope.h b/src/tir/transforms/update_pointer_storage_scope.h index d5e492e83389..5c082d24c407 100644 --- a/src/tir/transforms/update_pointer_storage_scope.h +++ b/src/tir/transforms/update_pointer_storage_scope.h @@ -39,10 +39,8 @@ class UpdatePointerStorageScope : public StmtExprMutator { const std::unordered_map& new_storage_scopes); virtual PrimExpr VisitExpr_(const VarNode*); - virtual PrimExpr VisitExpr_(const LoadNode*); virtual PrimExpr VisitExpr_(const BufferLoadNode*); virtual Stmt VisitStmt_(const AllocateNode*); - virtual Stmt VisitStmt_(const StoreNode*); virtual Stmt VisitStmt_(const BufferStoreNode*); private: diff --git a/src/tir/transforms/vectorize_loop.cc b/src/tir/transforms/vectorize_loop.cc index 6888ac625389..f809c07da021 100644 --- a/src/tir/transforms/vectorize_loop.cc +++ b/src/tir/transforms/vectorize_loop.cc @@ -63,14 +63,6 @@ class VecAllocAccess : public StmtExprMutator { VecAllocAccess(const VarNode* buf, Var var, int var_lanes) : buf_(buf), var_(var), var_lanes_(var_lanes) {} - PrimExpr VisitExpr_(const LoadNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated LoadNode. Please use BufferLoadNode instead."; - } - - Stmt VisitStmt_(const StoreNode* op) final { - LOG(FATAL) << "Unexpected use of deprecated StoreNode. Please use BufferStoreNode instead."; - } - PrimExpr VisitExpr_(const BufferLoadNode* op) final { auto load = Downcast(StmtExprMutator::VisitExpr_(op)); return UpdateBufferAccess(load); @@ -367,10 +359,6 @@ class Vectorizer : public StmtMutator, public ExprFunctor(op); @@ -414,10 +402,6 @@ class Vectorizer : public StmtMutator, public ExprFunctor(op); diff --git a/src/tir/usmp/analysis/extract_buffer_info.cc b/src/tir/usmp/analysis/extract_buffer_info.cc index 268058945750..f512bfaffa97 100644 --- a/src/tir/usmp/analysis/extract_buffer_info.cc +++ b/src/tir/usmp/analysis/extract_buffer_info.cc @@ -454,12 +454,7 @@ void BufferInfoExtractor::UpdateAliases(const Array& args, const PrimF // If tir.allocates are passed in to functions // The function params are re-directed to point // to the original allocate - if (arg->IsInstance()) { - auto load = Downcast(arg); - if (allocate_infos.count(load->buffer_var)) { - allocate_infos[param_buf] = allocate_infos[load->buffer_var]; - } - } else if (arg->IsInstance()) { + if (arg->IsInstance()) { auto var = Downcast(arg); if (allocate_infos.count(var)) { allocate_infos[param_buf] = allocate_infos[var]; diff --git a/src/tir/usmp/transform/create_io_allocates.cc b/src/tir/usmp/transform/create_io_allocates.cc index cf754131776c..0afdacd48fd7 100644 --- a/src/tir/usmp/transform/create_io_allocates.cc +++ b/src/tir/usmp/transform/create_io_allocates.cc @@ -52,10 +52,8 @@ class IOAllocateCreator : public StmtExprVisitor { private: void VisitExpr_(const BufferLoadNode* op) override; - void VisitExpr_(const LoadNode* op) override; void VisitExpr_(const CallNode* op) override; void VisitStmt_(const BufferStoreNode* op) override; - void VisitStmt_(const StoreNode* op) override; /*! \brief Updates aliases that buffer vars inside the primfunc refer * to in terms call arguments they get bound to.*/ @@ -150,8 +148,6 @@ void IOAllocateCreator::VisitExpr_(const BufferLoadNode* op) { StmtExprVisitor::VisitExpr_(op); } -void IOAllocateCreator::VisitExpr_(const LoadNode* op) { LOG(FATAL) << "should not come here"; } - void IOAllocateCreator::VisitStmt_(const BufferStoreNode* op) { if (aliases_.find(op->buffer->data) != aliases_.end()) { Var aliased_var = aliases_[op->buffer->data]; @@ -164,8 +160,6 @@ void IOAllocateCreator::VisitStmt_(const BufferStoreNode* op) { StmtExprVisitor::VisitStmt_(op); } -void IOAllocateCreator::VisitStmt_(const StoreNode* op) { LOG(FATAL) << "should not come here"; } - IRModule IOAllocateCreator::operator()() { Array new_main_params; Stmt main_body = main_func_->body; diff --git a/tests/python/integration/test_reduce.py b/tests/python/integration/test_reduce.py index 283eab3eea4c..f173e69cb94e 100644 --- a/tests/python/integration/test_reduce.py +++ b/tests/python/integration/test_reduce.py @@ -658,8 +658,8 @@ def run_passes(sch, args): # ... def check_store_dst_remapped(op): - if isinstance(op, tvm.tir.Store): - assert op.buffer_var.name != "reduce_temp0" + if isinstance(op, tvm.tir.BufferStore): + assert op.buffer.data.name != "reduce_temp0" tvm.tir.stmt_functor.post_order_visit(mod["main"].body, check_store_dst_remapped) diff --git a/tests/python/unittest/test_tir_transform_storage_rewrite.py b/tests/python/unittest/test_tir_transform_storage_rewrite.py index c46754fb1742..bcf498659902 100644 --- a/tests/python/unittest/test_tir_transform_storage_rewrite.py +++ b/tests/python/unittest/test_tir_transform_storage_rewrite.py @@ -242,11 +242,11 @@ def verify(v): # find add op if ( isinstance(v, tvm.tir.Add) - and isinstance(v.a, tvm.tir.Load) - and isinstance(v.b, tvm.tir.Load) + and isinstance(v.a, tvm.tir.BufferLoad) + and isinstance(v.b, tvm.tir.BufferLoad) ): - lhs_ramp = v.a.index - rhs_ramp = v.b.index + lhs_ramp = v.a.indices[0] + rhs_ramp = v.b.indices[0] # these two ramp load should not overlap assert lhs_ramp.lanes == n assert rhs_ramp.lanes == n