diff --git a/include/tvm/tir/buffer.h b/include/tvm/tir/buffer.h index aef82ae368d0..876a701419ea 100644 --- a/include/tvm/tir/buffer.h +++ b/include/tvm/tir/buffer.h @@ -76,18 +76,33 @@ class BufferNode : public Object { * This can be an empty array, indicating array is contiguous */ Array strides; - /*! \brief The offset in terms of number of dtype elements (including lanes) */ - PrimExpr elem_offset; + + /*! \brief The offset for each physical dimension of the buffer. + * + * This is the offset of the start of the buffer, relative to the + * `data` variable. This is typically used to represent offsets + * between two buffers that are backed by the same allocation, such + * as a buffer bind in a tensorized computation. The offset is in + * terms of the scalar type of the buffer (e.g. an offset of 8 in a + * buffer of dtype float16x4 is a byte offset of 16 (8 elements) * + * (2 bytes/float16)). + */ + Array elem_offsets; + // Meta data /*! \brief optional name of the buffer */ String name; /*! \brief Alignment requirement of data pointer in bytes. */ int data_alignment; + /*! - * \brief Factor of elem_offset field, - * elem_offset is guaranteed to be multiple of offset_factor. + * \brief Factor of elem_offset field. + * + * Each value in `elem_offsets` is guaranteed to be multiple of the + * corresponding element in `offset_factors`. */ - int offset_factor; + Array offset_factors; + /*! \brief buffer type */ BufferType buffer_type; /*! @@ -104,10 +119,10 @@ class BufferNode : public Object { v->Visit("shape", &shape); v->Visit("strides", &strides); v->Visit("axis_separators", &axis_separators); - v->Visit("elem_offset", &elem_offset); + v->Visit("elem_offsets", &elem_offsets); v->Visit("name", &name); v->Visit("data_alignment", &data_alignment); - v->Visit("offset_factor", &offset_factor); + v->Visit("offset_factors", &offset_factors); v->Visit("buffer_type", &buffer_type); v->Visit("span", &span); } @@ -118,7 +133,7 @@ class BufferNode : public Object { return equal.DefEqual(data, other->data) && equal(dtype, other->dtype) && equal.DefEqual(shape, other->shape) && equal.DefEqual(strides, other->strides) && equal.DefEqual(axis_separators, other->axis_separators) && - equal.DefEqual(elem_offset, other->elem_offset) && + equal.DefEqual(elem_offsets, other->elem_offsets) && equal(data_alignment, other->data_alignment) && equal(buffer_type, other->buffer_type); } @@ -127,16 +142,14 @@ class BufferNode : public Object { hash_reduce(dtype); hash_reduce.DefHash(shape); hash_reduce.DefHash(strides); - hash_reduce.DefHash(elem_offset); + hash_reduce.DefHash(elem_offsets); hash_reduce.DefHash(axis_separators); hash_reduce(data_alignment); hash_reduce(buffer_type); } /*! \return preferred index type for this buffer node */ - DataType DefaultIndexType() const { - return shape.size() != 0 ? shape[0].dtype() : DataType::Int(32); - } + DataType DefaultIndexType() const { return DefaultIndexType(shape); } /*! \brief Determine the offset in the buffer of the given index. * @@ -150,6 +163,12 @@ class BufferNode : public Object { static constexpr const bool _type_has_method_sequal_reduce = true; static constexpr const bool _type_has_method_shash_reduce = true; TVM_DECLARE_FINAL_OBJECT_INFO(BufferNode, Object); + + private: + static DataType DefaultIndexType(const Array& shape) { + return shape.size() != 0 ? shape[0].dtype() : DataType::Int(32); + } + friend class Buffer; }; /*! @@ -165,6 +184,13 @@ class Buffer : public ObjectRef { PrimExpr elem_offset, String name, int data_alignment, int offset_factor, BufferType buffer_type, Array axis_separators = {}, Span span = Span()); + // User can specify data_alignment and offset_factor to be 0 + // A default value will be picked. + TVM_DLL Buffer(Var data, DataType dtype, Array shape, Array strides, + Array elem_offsets, String name, int data_alignment, + Array offset_factors, BufferType buffer_type, + Array axis_separators = {}, Span span = Span()); + /*! * \brief Return a new buffer that is equivalent with current one * but always add stride field. diff --git a/include/tvm/topi/detail/extern.h b/include/tvm/topi/detail/extern.h index 2561f8d1ca27..e715e9f5ad87 100644 --- a/include/tvm/topi/detail/extern.h +++ b/include/tvm/topi/detail/extern.h @@ -122,12 +122,19 @@ inline PrimExpr pack_buffer(Buffer buf) { } else { strides = 0; } + + CHECK(buf->axis_separators.empty()) + << "Cannot represent " << buf->axis_separators.size() + 1 << "-d buffer as a DLTensor"; + ICHECK_LE(buf->elem_offsets.size(), 1); + + PrimExpr elem_offset = buf->elem_offsets.size() ? buf->elem_offsets[0] : 0; + Array pack_args{buf->data, shape, strides, make_const(DataType::Int(32), static_cast(buf->shape.size())), make_const(buf->dtype, 0), - buf->elem_offset}; + elem_offset}; return tvm::tir::Call(DataType::Handle(), tvm::tir::builtin::tvm_stack_make_array(), pack_args); } diff --git a/python/tvm/script/tir/special_stmt.py b/python/tvm/script/tir/special_stmt.py index 0148bd0b4243..c01332f247c1 100644 --- a/python/tvm/script/tir/special_stmt.py +++ b/python/tvm/script/tir/special_stmt.py @@ -143,9 +143,19 @@ def match_buffer( if strides is None: strides = [] align = convert_to_int(align, "align", self.context.report_error, self.node.span) - offset_factor = convert_to_int( - offset_factor, "offset_factor", self.context.report_error, self.node.span - ) + + if offset_factor != 0: + try: + offset_factor = list(offset_factor) + except TypeError: + offset_factor = [offset_factor] + + offset_factor = [ + convert_to_int( + factor, "offset_factor", self.context.report_error, self.node.span + ) + for factor in offset_factor + ] buffer_name: str = self.node.lhs[0].id.name buffer = tvm.tir.decl_buffer( shape, diff --git a/python/tvm/tir/buffer.py b/python/tvm/tir/buffer.py index e36a99339e48..af9286e269a0 100644 --- a/python/tvm/tir/buffer.py +++ b/python/tvm/tir/buffer.py @@ -170,6 +170,54 @@ def offset_of(self, indices): """ return _ffi_api.BufferOffsetOf(self, indices) # type: ignore + @property + def elem_offset(self): + """The offset of the buffer, relative to the backing + allocation. + + This is provided for backwards compatibility for flat memory + spaces. For non-flat memory spaces, each physical dimension + has an independent offset, which can be accessed with + `buf.elem_offsets`. + + Returns + ------- + elem_offset: PrimExpr + + The offset of the buffer, relative to the backing + allocation. + """ + + offsets = self.elem_offsets + assert ( + len(offsets) == 1 + ), "Non-unique buffer.elem_offset, use buffer.elem_offsets (plural) instead." + return offsets[0] + + @property + def offset_factor(self): + """the offset of the buffer, relative to the backing + allocation. + + This is provided for backwards compatibility for flat memory + spaces. For non-flat memory spaces, each physical dimension + has an independent offset factor, which can be accessed with + `buf.offset_factors`. + + Returns + ------- + offset_factor: int + + The offset of the buffer, relative to the backing + allocation. + """ + + factors = self.offset_factors + assert ( + len(factors) == 1 + ), "Non-unique buffer.offset_factor, use buffer.offset_factors (plural) instead." + return factors[0] + def decl_buffer( shape, @@ -209,9 +257,18 @@ def decl_buffer( strides: array of Expr The stride of the buffer. - elem_offset: Expr, optional - The beginning offset of the array to data. - In terms of number of elements of dtype. + elem_offset: Optional[Union[Expr, Sequence[Expr]]] + + The beginning offset of the array to data, in terms of number + of elements of dtype. If None, will be treated either as an + offset of zero for each physical dimension if offset_factor is + unset, or as a list of Var objects that are implicitly defined + by this buffer declaration. If an expression is passed, is + equivalent to passing a list of length 1. + + The number of offsets should be the number of physical + dimensions of the buffer, or one more than the number of + `axis_separators`. scope: str, optional The storage scope of the buffer, if not global. @@ -221,11 +278,14 @@ def decl_buffer( The alignment of data pointer in bytes. If -1 is passed, the alignment will be set to TVM's internal default. - offset_factor: int, optional - The factor of elem_offset field, when set, - elem_offset is required to be multiple of offset_factor. - If 0 is pssed, the alignment will be set to 1. - if non-zero is passed, we will created a Var for elem_offset if elem_offset is not None. + offset_factor: Optional[Union[int, Sequence[int]]] + + The factor of elem_offset field. When set, elem_offset is + required to be multiple of offset_factor. If 0 is passed, the + alignment will be set to 1 on all physical dimensions. If + non-zero is passed and elem_offset is None, elem_offset will + be set to a list of Var objects representing their implicit + definition by this buffer declaration. buffer_type: str, optional, {"", "auto_broadcast"} auto_broadcast buffer allows one to implement broadcast computation @@ -280,6 +340,7 @@ def decl_buffer( for the DLTensor that is compact and aligned. If user pass a fully generic symbolic array to the strides, then the resulting function becomes fully generic. + """ # pylint: disable=import-outside-toplevel from .expr import Var @@ -291,14 +352,31 @@ def decl_buffer( if axis_separators is None: axis_separators = [] - if offset_factor != 0 and elem_offset is None: + n_physical_dim = len(axis_separators) + 1 + + if offset_factor == 0 or offset_factor is None: + offset_factor = [0 for _ in range(n_physical_dim)] + elif isinstance(offset_factor, int): + offset_factor = [offset_factor] + + if elem_offset is None: shape_dtype = shape[0].dtype if shape and hasattr(shape[0], "dtype") else "int32" - elem_offset = Var("%s_elem_offset" % name, shape_dtype) + elem_offset = [ + 0 if factor == 0 else Var(f"{name}_elem_offset_{i}", shape_dtype) + for i, factor in enumerate(offset_factor) + ] + + try: + elem_offset = list(elem_offset) + except TypeError: + elem_offset = [elem_offset] + if data is None: # Bool is represented as uint1 in the IR, but stored as int8 storage_type = PrimType(dtype) storage_type = PrimType("int8") if storage_type.dtype == "bool" else storage_type data = Var(name, PointerType(storage_type, scope), span) + return _ffi_api.Buffer( # type: ignore data, dtype, diff --git a/src/printer/tir_text_printer.cc b/src/printer/tir_text_printer.cc index 16d477232eb2..191cd610d815 100644 --- a/src/printer/tir_text_printer.cc +++ b/src/printer/tir_text_printer.cc @@ -220,8 +220,8 @@ Doc TIRTextPrinter::PrintProducer(const DataProducerNode* op) { Doc TIRTextPrinter::BufferNode2Doc(const BufferNode* buf, Doc doc) { doc << Doc::Text(": Buffer(") << Print(buf->data) << ", " << PrintDType(buf->dtype) << ", " << Print(buf->shape) << ", " << Print(buf->strides); - if (!is_zero(buf->elem_offset)) { - doc << ", elem_offset=" << Print(buf->elem_offset); + if (buf->elem_offsets.size()) { + doc << ", elem_offsets=" << Print(buf->elem_offsets); } if (buf->axis_separators.size()) { doc << ", axis_separators=" << Print(buf->axis_separators); @@ -232,9 +232,18 @@ Doc TIRTextPrinter::BufferNode2Doc(const BufferNode* buf, Doc doc) { if (buf->data_alignment != 128) { doc << ", align=" << buf->data_alignment; } - if (buf->offset_factor != 1) { - doc << ", offset_factor=" << buf->offset_factor; + + bool has_non_default_offset_factor = false; + for (const auto& offset_factor : buf->offset_factors) { + if (offset_factor->value != 1) { + has_non_default_offset_factor = true; + break; + } } + if (has_non_default_offset_factor) { + doc << ", offset_factor=" << Print(buf->offset_factors); + } + if (buf->buffer_type != 1) { doc << ", type=" << Doc::StrLiteral("auto"); } diff --git a/src/printer/tvmscript_printer.cc b/src/printer/tvmscript_printer.cc index da5975cd5e28..50f7fc0a505d 100644 --- a/src/printer/tvmscript_printer.cc +++ b/src/printer/tvmscript_printer.cc @@ -474,30 +474,73 @@ Doc TVMScriptPrinter::AllocBufferDeclaration(const Buffer& buf) { if (!buf->strides.empty()) { doc << ", strides=" << Print(buf->strides); } - if (buf->elem_offset->IsInstance()) { - Var elem_offset = Downcast(buf->elem_offset); - if (memo_var_.find(elem_offset) != memo_var_.end()) { - doc << ", elem_offset=" << Print(buf->elem_offset); - } else { - // implicitly define elem_offset - memo_var_[elem_offset] = Doc::Text(memo_buf_[buf].str() + ".elem_offset"); - var_not_in_headers_.insert(elem_offset.get()); - print_factor_explicitly = true; + + { + bool requires_elem_offset_print = false; + bool is_implicit_definition = false; + bool is_explicit_usage = false; + + Doc doc_offsets; + + doc_offsets << "["; + for (size_t i = 0; i < buf->elem_offsets.size(); i++) { + if (i) { + doc_offsets << ", "; + } + + auto offset = buf->elem_offsets[i]; + if (offset->IsInstance() && + memo_var_.find(Downcast(offset)) == memo_var_.end()) { + // If the element offset is a Var not previously encountered, then + // the buffer declaration is an implicit definition for that + // variable. + CHECK(!is_explicit_usage) + << "Cannot mix implicit definitions and explict variable use in elem_offsets"; + + requires_elem_offset_print = false; + is_implicit_definition = true; + std::stringstream ss; + ss << memo_buf_[buf].str() << ".elem_offsets[" << i << "]"; + memo_var_[Downcast(offset)] = Doc::Text(ss.str()); + + } else { + // Otherwise, this is an argument that must be passed to the + // buffer declaration, and should be printed. + CHECK(!is_implicit_definition) + << "Cannot mix implicit definitions and explict variable use in elem_offsets"; + is_explicit_usage = true; + + doc_offsets << Print(offset); + + auto int_offset = offset.as(); + if (!int_offset || int_offset->value != 0) { + requires_elem_offset_print = true; + } + } } - } else if (buf->elem_offset->IsInstance()) { - IntImm elem_offset = Downcast(buf->elem_offset); - if (elem_offset->value != 0) { - doc << ", elem_offset=" << Print(buf->elem_offset); + doc_offsets << "]"; + + if (requires_elem_offset_print) { + doc << ", elem_offsets=" << doc_offsets; } } + if (buf.scope() != "global") { doc << ", scope=" << Doc::StrLiteral(buf.scope()); } if (buf->data_alignment != runtime::kAllocAlignment) { doc << ", align=" << buf->data_alignment; } - if (buf->offset_factor != 1 || print_factor_explicitly) { - doc << ", offset_factor=" << buf->offset_factor; + + bool has_non_default_offset_factor = false; + for (const auto& offset_factor : buf->offset_factors) { + if (offset_factor->value != 1) { + has_non_default_offset_factor = true; + break; + } + } + if (has_non_default_offset_factor || print_factor_explicitly) { + doc << ", offset_factors=" << Print(buf->offset_factors); } if (buf->buffer_type != BufferType::kDefault) { doc << ", type=" << Doc::StrLiteral("auto"); @@ -586,12 +629,13 @@ bool TVMScriptPrinter::IsSimpleBuffer(const Buffer& buf) { return false; } } - if (!UndefinedVars(buf->elem_offset).empty()) { - return false; - } else if (buf->elem_offset->IsInstance()) { - IntImm elem_offset = Downcast(buf->elem_offset); - if (elem_offset->value != 0) { + for (const auto& offset : buf->elem_offsets) { + if (!UndefinedVars(offset).empty()) { return false; + } else if (auto* int_offset = offset.as()) { + if (int_offset->value != 0) { + return false; + } } } if (buf.scope() != "global") { @@ -600,8 +644,11 @@ bool TVMScriptPrinter::IsSimpleBuffer(const Buffer& buf) { if (buf->data_alignment != runtime::kAllocAlignment) { return false; } - if (buf->offset_factor != 1) { - return false; + + for (const auto& offset_factor : buf->offset_factors) { + if (offset_factor->value != 1) { + return false; + } } if (buf->buffer_type != BufferType::kDefault) { return false; diff --git a/src/te/schedule/schedule_postproc_to_primfunc.cc b/src/te/schedule/schedule_postproc_to_primfunc.cc index c7d5d7a8dafa..f89cecc19d53 100644 --- a/src/te/schedule/schedule_postproc_to_primfunc.cc +++ b/src/te/schedule/schedule_postproc_to_primfunc.cc @@ -337,8 +337,9 @@ class AxisSeparatorsAttrUnwrapper : StmtExprMutator { if (lookup) { Array axis_separators = lookup.value(); if (axis_separators.size()) { - auto write_ptr = buf.CopyOnWrite(); - write_ptr->axis_separators = axis_separators; + buf = Buffer(buf->data, buf->dtype, buf->shape, buf->strides, buf->elem_offsets, buf->name, + buf->data_alignment, buf->offset_factors, buf->buffer_type, + buf->axis_separators, buf->span); } } diff --git a/src/tir/analysis/verify_ssa.cc b/src/tir/analysis/verify_ssa.cc index d7ccb363c16e..2e066890e73b 100644 --- a/src/tir/analysis/verify_ssa.cc +++ b/src/tir/analysis/verify_ssa.cc @@ -100,16 +100,19 @@ class SSAVerifier final : public StmtExprVisitor { void DefineBuffer(const Buffer& buffer) { match_scope_ = true; this->VisitExpr(buffer->data); - for (size_t i = 0; i < buffer->shape.size(); ++i) { - this->VisitExpr(buffer->shape[i]); + for (const auto& dim : buffer->shape) { + this->VisitExpr(dim); } if (buffer->strides.defined()) { - for (size_t i = 0; i < buffer->strides.size(); ++i) { - this->VisitExpr(buffer->strides[i]); + for (const auto& stride : buffer->strides) { + this->VisitExpr(stride); } } - this->VisitExpr(buffer->elem_offset); + + for (const auto& offset : buffer->elem_offsets) { + this->VisitExpr(offset); + } match_scope_ = false; } diff --git a/src/tir/ir/buffer.cc b/src/tir/ir/buffer.cc index 4fe9b162078e..f042375c6726 100644 --- a/src/tir/ir/buffer.cc +++ b/src/tir/ir/buffer.cc @@ -262,14 +262,6 @@ Array BufferNode::ElemOffset(Array input_indices) const { << "the index's dimensionality must match the dimensionality of the index given."; } - // TODO(Lunderberg): Better handling for cases where there is more - // than one output index. Currently, this only allows elem_offset - // to be non-zero for flat memory allocations. - Array elem_offsets = {}; - if (elem_offset.defined() && !is_zero(elem_offset)) { - elem_offsets = {elem_offset}; - } - if (elem_offsets.size()) { ICHECK_EQ(elem_offsets.size(), axis_separators.size() + 1) << "If element offsets are defined, " @@ -488,6 +480,12 @@ PrimExpr Buffer::access_ptr(int access_mask, DataType ptr_type, int content_lane PrimExpr offset) const { const BufferNode* self = operator->(); ICHECK(self != nullptr); + + // To relax this restriction, will need to have `tvm_access_ptr` + // keep an array of offsets. Maybe best to hold the type, buffer + // var, and offset all bundled together into a BufferLoad? + ICHECK(self->axis_separators.empty()) << "access_ptr not yet supported on non-flat memory"; + PrimExpr e_dtype; PrimExpr extent; if (self->shape.size() == 0) { @@ -500,11 +498,14 @@ PrimExpr Buffer::access_ptr(int access_mask, DataType ptr_type, int content_lane make_const(DataType::Int(32), 1), self->shape) - offset; } - PrimExpr elem_offset = self->elem_offset + offset; + + ICHECK_EQ(self->elem_offsets.size(), 1); + + PrimExpr elem_offset = self->elem_offsets[0] + offset; if (content_lanes > 1) { e_dtype = tir::TypeAnnotation(self->dtype.with_lanes(content_lanes)); - extent = extent / make_const(self->elem_offset.dtype(), content_lanes); - elem_offset = self->elem_offset / make_const(self->elem_offset.dtype(), content_lanes); + extent = extent / make_const(self->elem_offsets[0].dtype(), content_lanes); + elem_offset = self->elem_offsets[0] / make_const(self->elem_offsets[0].dtype(), content_lanes); } else { e_dtype = tir::TypeAnnotation(self->dtype); } @@ -515,7 +516,15 @@ PrimExpr Buffer::access_ptr(int access_mask, DataType ptr_type, int content_lane Buffer::Buffer(Var data, DataType dtype, Array shape, Array strides, PrimExpr elem_offset, String name, int data_alignment, int offset_factor, - BufferType buffer_type, Array axis_separators, Span span) { + BufferType buffer_type, Array axis_separators, Span span) + : Buffer(data, dtype, shape, strides, {elem_offset}, name, data_alignment, + Array{IntImm(BufferNode::DefaultIndexType(shape), offset_factor)}, buffer_type, + axis_separators, span) {} + +Buffer::Buffer(Var data, DataType dtype, Array shape, Array strides, + Array elem_offsets, String name, int data_alignment, + Array offset_factors, BufferType buffer_type, Array axis_separators, + Span span) { DataType storage_dtype = dtype; // specially handle bool if (storage_dtype == DataType::Bool()) { @@ -537,6 +546,36 @@ Buffer::Buffer(Var data, DataType dtype, Array shape, Array ICHECK(data->type_annotation.as()->element_type.as()) << "Variable " << data->name_hint << " does not point to a primitive."; + size_t n_physical_dim = axis_separators.size() + 1; + + DataType index_dtype = BufferNode::DefaultIndexType(shape); + for (size_t i = 0; i < elem_offsets.size(); i++) { + if (!elem_offsets[i].defined()) { + elem_offsets.Set(i, make_const(index_dtype, 0)); + } + } + if (elem_offsets.size() == 0 || (elem_offsets.size() == 1 && is_zero(elem_offsets[0]))) { + elem_offsets = Array(n_physical_dim, IntImm(index_dtype, 0)); + } + + if (data_alignment <= 0) { + data_alignment = runtime::kAllocAlignment; + } + + for (size_t i = 0; i < offset_factors.size(); i++) { + if (offset_factors[i]->value == 0) { + offset_factors.Set(i, IntImm(index_dtype, 1)); + } + } + if (offset_factors.size() == 0) { + offset_factors = Array(n_physical_dim, IntImm(index_dtype, 1)); + } + + CHECK_EQ(elem_offsets.size(), n_physical_dim) + << "Expected one element offset for each physical dimension of the buffer"; + CHECK_EQ(offset_factors.size(), n_physical_dim) + << "Expected one offset factor for each physical dimension of the buffer"; + auto n = make_object(); n->data = std::move(data); n->dtype = dtype; @@ -545,18 +584,9 @@ Buffer::Buffer(Var data, DataType dtype, Array shape, Array n->strides = std::move(strides); n->axis_separators = std::move(axis_separators); n->name = std::move(name); - if (!elem_offset.defined()) { - elem_offset = make_const(n->DefaultIndexType(), 0); - } - if (data_alignment <= 0) { - data_alignment = runtime::kAllocAlignment; - } - if (offset_factor == 0) { - offset_factor = 1; - } - n->elem_offset = std::move(elem_offset); + n->elem_offsets = std::move(elem_offsets); n->data_alignment = data_alignment; - n->offset_factor = offset_factor; + n->offset_factors = offset_factors; n->buffer_type = buffer_type; if (n->buffer_type == kAutoBroadcast && n->shape.size() > 0 && n->strides.empty()) { for (size_t i = 0; i < n->shape.size(); ++i) { @@ -577,9 +607,12 @@ TVM_REGISTER_NODE_TYPE(BufferNode); TVM_REGISTER_GLOBAL("tir.Buffer").set_body([](TVMArgs args, TVMRetValue* ret) { ICHECK_EQ(args.size(), 11); + // Resolve overloaded Buffer constructor by passing a typed argument + // for the element offsets. + Array elem_offsets = args[4].operator Array(); auto buffer_type = args[8].operator String(); BufferType type = (buffer_type == "auto_broadcast") ? kAutoBroadcast : kDefault; - *ret = Buffer(args[0], args[1], args[2], args[3], args[4], args[5], args[6], args[7], type, + *ret = Buffer(args[0], args[1], args[2], args[3], elem_offsets, args[5], args[6], args[7], type, args[9], args[10]); }); diff --git a/src/tir/ir/specialize.cc b/src/tir/ir/specialize.cc index 1e5b2f28b2d9..6a727d4a8a28 100644 --- a/src/tir/ir/specialize.cc +++ b/src/tir/ir/specialize.cc @@ -200,19 +200,17 @@ class PrimFuncSpecializer : public StmtExprMutator { private: Buffer MutateBuffer(const Buffer& buffer) { - Array shape = - MutateArray(buffer->shape, [this](const PrimExpr& e) { return VisitExpr(e); }); - Array strides = - MutateArray(buffer->strides, [this](const PrimExpr& e) { return VisitExpr(e); }); + auto mutate = [this](const PrimExpr& e) { return VisitExpr(e); }; + Array shape = MutateArray(buffer->shape, mutate); + Array strides = MutateArray(buffer->strides, mutate); + Array elem_offsets = MutateArray(buffer->elem_offsets, mutate); - PrimExpr elem_offset = VisitExpr(buffer->elem_offset); - - if (buffer->elem_offset.same_as(elem_offset) && buffer->shape.same_as(shape) && + if (buffer->elem_offsets.same_as(elem_offsets) && buffer->shape.same_as(shape) && buffer->strides.same_as(strides)) { return buffer; } else { auto n = make_object(*buffer.get()); - n->elem_offset = std::move(elem_offset); + n->elem_offsets = std::move(elem_offsets); n->shape = std::move(shape); n->strides = std::move(strides); return Buffer(n); @@ -313,6 +311,11 @@ void UpdateSpecializeVarMap(const PrimFunc& func, const Var& param, const Buffer << "ValueError: The buffer strides dimensions mismatched" << buf_to_specialize->strides.size() << " vs. " << specific_buf->strides.size() << "."; + CHECK(specific_buf->elem_offsets.size() == buf_to_specialize->elem_offsets.size()) + << "ValueError: The buffer offsets dimensions mismatched" + << buf_to_specialize->elem_offsets.size() << " vs. " << specific_buf->elem_offsets.size() + << "."; + // Updating var mapping using specific_expr for (size_t i = 0; i < specific_buf->shape.size(); ++i) { build_var_mapping(specific_buf->shape[i], buf_to_specialize->shape[i]); @@ -320,7 +323,9 @@ void UpdateSpecializeVarMap(const PrimFunc& func, const Var& param, const Buffer for (size_t i = 0; i < specific_buf->strides.size(); ++i) { build_var_mapping(specific_buf->strides[i], buf_to_specialize->strides[i]); } - build_var_mapping(specific_buf->elem_offset, buf_to_specialize->elem_offset); + for (size_t i = 0; i < specific_buf->elem_offsets.size(); ++i) { + build_var_mapping(specific_buf->elem_offsets[i], buf_to_specialize->elem_offsets[i]); + } // Check data_alignment and offset_factor. // These two signatures are int, so we do not need map them. @@ -328,9 +333,17 @@ void UpdateSpecializeVarMap(const PrimFunc& func, const Var& param, const Buffer << "ValueError: The buffer data_alignment mismatched" << buf_to_specialize->data_alignment << " vs. " << specific_buf->data_alignment << "."; - CHECK_EQ(specific_buf->offset_factor, buf_to_specialize->offset_factor) - << "ValueError: The buffer offset_factor mismatched" << buf_to_specialize->offset_factor - << " vs. " << specific_buf->offset_factor << "."; + CHECK(specific_buf->offset_factors.size() == buf_to_specialize->offset_factors.size()) + << "ValueError: The buffer offset factors dimensions mismatched" + << buf_to_specialize->offset_factors.size() << " vs. " << specific_buf->offset_factors.size() + << "."; + + for (size_t i = 0; i < specific_buf->offset_factors.size(); ++i) { + CHECK_EQ(specific_buf->offset_factors[i]->value, buf_to_specialize->offset_factors[i]->value) + << "ValueError: The buffer offset_factors[" << i << "] mismatched" + << buf_to_specialize->offset_factors[i] << " vs. " << specific_buf->offset_factors[i] + << "."; + } } /*! diff --git a/src/tir/transforms/arg_binder.cc b/src/tir/transforms/arg_binder.cc index d7cd731a3d2b..ae137271eec8 100644 --- a/src/tir/transforms/arg_binder.cc +++ b/src/tir/transforms/arg_binder.cc @@ -97,21 +97,35 @@ void ArgBinder::BindBuffer(const Buffer& arg, const Buffer& value, const std::st << ", provided_alignment=" << value->data_alignment; } // bind pointer and offset. - if (is_zero(arg->elem_offset)) { - ICHECK(is_zero(value->elem_offset)) - << "Trying to bind a Buffer with offset into one without offset " - << " required elem_offset=" << arg->elem_offset - << ", provided elem_offset=" << value->elem_offset; + ICHECK_EQ(arg->elem_offsets.size(), value->elem_offsets.size()) + << "Trying to bind buffer with different physical dimension, requires " + << arg->elem_offsets.size() << "-d buffer, but provided " << value->elem_offsets.size() + << "-d buffer"; + for (size_t i = 0; i < arg->elem_offsets.size(); i++) { + auto arg_offset = arg->elem_offsets[i]; + if (is_zero(arg_offset)) { + auto value_offset = value->elem_offsets[i]; + ICHECK(is_zero(value_offset)) + << "Trying to bind a Buffer with offset into one without offset " + << " required elem_offset=" << arg_offset << ", provided elem_offset=" << value_offset; + } } this->Bind(arg->data, value->data, arg_name + ".data"); - if (Bind_(arg->elem_offset, value->elem_offset, arg_name + ".elem_offset", false)) { - if (arg->offset_factor > 1) { - PrimExpr offset = value->elem_offset; - PrimExpr factor = make_const(offset.dtype(), arg->offset_factor); - PrimExpr zero = make_zero(offset.dtype()); - BinderAddAssert(&analyzer_, truncmod(offset, factor) == zero, arg_name + ".elem_offset", - &asserts_); + + ICHECK_EQ(arg->elem_offsets.size(), arg->offset_factors.size()); + ICHECK_EQ(value->elem_offsets.size(), value->offset_factors.size()); + for (size_t i = 0; i < arg->elem_offsets.size(); i++) { + auto arg_offset = arg->elem_offsets[i]; + auto value_offset = value->elem_offsets[i]; + if (Bind_(arg_offset, value_offset, arg_name + ".elem_offset", false)) { + auto arg_offset_factor = arg->offset_factors[i]->value; + if (arg_offset_factor > 1) { + PrimExpr factor = make_const(value_offset.dtype(), arg_offset_factor); + PrimExpr zero = make_zero(value_offset.dtype()); + BinderAddAssert(&analyzer_, truncmod(value_offset, factor) == zero, + arg_name + ".elem_offset", &asserts_); + } } } @@ -271,24 +285,36 @@ void ArgBinder::BindDLTensor(const Buffer& buffer, const PrimExpr& device_type, // Byte_offset field. int data_bytes = GetVectorBytes(buffer->dtype); - if (const auto* const_offset = buffer->elem_offset.as()) { - Bind_(make_const(DataType::UInt(64), const_offset->value * data_bytes), - TVMArrayGet(DataType::UInt(64), handle, builtin::kArrByteOffset), - arg_name + ".byte_offset", true); - } else { - if (Bind_(buffer->elem_offset, - cast(buffer->elem_offset.dtype(), - (TVMArrayGet(DataType::UInt(64), handle, builtin::kArrByteOffset) / - make_const(DataType::UInt(64), data_bytes))), + PrimExpr arg_byte_offset = TVMArrayGet(DataType::UInt(64), handle, builtin::kArrByteOffset); + if (buffer->elem_offsets.size() == 1) { + auto offset = buffer->elem_offsets[0]; + + if (const auto* const_offset = offset.as()) { + Bind_(make_const(DataType::UInt(64), const_offset->value * data_bytes), arg_byte_offset, + arg_name + ".byte_offset", true); + } else { + if (Bind_( + offset, + cast(offset.dtype(), (arg_byte_offset / make_const(DataType::UInt(64), data_bytes))), arg_name + ".elem_offset", true)) { - if (buffer->offset_factor > 1) { - PrimExpr offset = buffer->elem_offset; - PrimExpr factor = make_const(offset.dtype(), buffer->offset_factor); - PrimExpr zero = make_zero(offset.dtype()); - BinderAddAssert(&analyzer_, truncmod(offset, factor) == zero, arg_name + ".elem_offset", - &asserts_); + auto factor = buffer->offset_factors[0]; + if (factor->value > 1) { + PrimExpr zero = make_zero(offset.dtype()); + BinderAddAssert(&analyzer_, truncmod(offset, factor) == zero, arg_name + ".elem_offset", + &asserts_); + } } } + + } else { + for (size_t i = 0; i < buffer->elem_offsets.size(); i++) { + auto offset = buffer->elem_offsets[i]; + CHECK(!is_zero(offset)) << "Buffer " << buffer->name << ".elem_offsets[" << i + << "] = " << tvm::PrettyPrint(offset) + << ", but non-zero element offsets across function boundaries " + << "are only supported for flat memory spaces."; + } + BinderAddAssert(&analyzer_, arg_byte_offset == 0, arg_name + ".byte_offset", &asserts_); } // device info. Bind_(device_type, TVMArrayGet(DataType::Int(32), handle, builtin::kArrDeviceType), diff --git a/src/tir/transforms/arg_binder.h b/src/tir/transforms/arg_binder.h index 657ebdbec134..43ce9cc4dd95 100644 --- a/src/tir/transforms/arg_binder.h +++ b/src/tir/transforms/arg_binder.h @@ -122,7 +122,25 @@ class ArgBinder { const Map& def_handle_dtype() const { return def_handle_dtype_; } private: - // Internal bind function + /* \brief Internal bind function + * + * \param arg The expression that occurs within a function body. + * + * \param value The expression that whose value is bound to arg for + * the duration of the function. + * + * \param arg_name The name of the argument being bound. This is + * used for creating error messages. + * + * \param with_lets If true, variable definitions should be done by + * inserting a let statement with the function initialiation in + * `init_nest_`. If false, variable definitions should be + * performed by adding a substitution to `def_map_`. If `arg` + * is not a `VarNode`, `with_lets` is unused. + * + * \return True if arg is a VarNode being defined by this binding, + * otherwise false. + */ bool Bind_(const PrimExpr& arg, const PrimExpr& value, const std::string& arg_name, bool with_lets); /*! \brief The definition map, can be uses to substitute */ diff --git a/src/tir/transforms/bf16_legalize.cc b/src/tir/transforms/bf16_legalize.cc index 193584f84b47..811901c9ea51 100644 --- a/src/tir/transforms/bf16_legalize.cc +++ b/src/tir/transforms/bf16_legalize.cc @@ -276,9 +276,9 @@ class BF16LowerRewriter : public StmtExprMutator { if (oldbuf->dtype.is_bfloat16()) { DataType dtype = DataType::UInt(16, oldbuf->dtype.lanes()); Var buffer_var = Var(oldbuf->data->name_hint, PointerType(PrimType(dtype))); - auto newbuf = Buffer(buffer_var, dtype, oldbuf->shape, oldbuf->strides, oldbuf->elem_offset, - oldbuf->name, oldbuf->data_alignment, oldbuf->offset_factor, - oldbuf->buffer_type); + auto newbuf = Buffer(buffer_var, dtype, oldbuf->shape, oldbuf->strides, + oldbuf->elem_offsets, oldbuf->name, oldbuf->data_alignment, + oldbuf->offset_factors, oldbuf->buffer_type); buffer_remap_[oldbuf] = newbuf; var_remap_[oldbuf->data] = buffer_var; new_buffer_map.Set(param_var, newbuf); @@ -306,8 +306,8 @@ class BF16LowerRewriter : public StmtExprMutator { const Buffer& flatbuf = (*it).second; DataType dtype = DataType::UInt(16, oldbuf->dtype.lanes()); auto newbuf = Buffer(flatbuf->data, dtype, oldbuf->shape, oldbuf->strides, - oldbuf->elem_offset, oldbuf->name, oldbuf->data_alignment, - oldbuf->offset_factor, oldbuf->buffer_type); + oldbuf->elem_offsets, oldbuf->name, oldbuf->data_alignment, + oldbuf->offset_factors, oldbuf->buffer_type); buffer_remap_[oldbuf] = newbuf; new_preflattened_buffer_map.Set(param_var, newbuf); } else { @@ -334,8 +334,8 @@ class BF16LowerRewriter : public StmtExprMutator { if (var_it != var_remap_.end()) { DataType dtype = buf->dtype.is_bfloat16() ? DataType::UInt(16, buf->dtype.lanes()) : buf->dtype; - new_buf = Buffer(var_it->second, dtype, buf->shape, buf->strides, buf->elem_offset, buf->name, - buf->data_alignment, buf->offset_factor, buf->buffer_type, + new_buf = Buffer(var_it->second, dtype, buf->shape, buf->strides, buf->elem_offsets, + buf->name, buf->data_alignment, buf->offset_factors, buf->buffer_type, buf->axis_separators, buf->span); } diff --git a/src/tir/transforms/inject_copy_intrin.cc b/src/tir/transforms/inject_copy_intrin.cc index 81842ff808ff..5ca1d07449c2 100644 --- a/src/tir/transforms/inject_copy_intrin.cc +++ b/src/tir/transforms/inject_copy_intrin.cc @@ -174,7 +174,7 @@ class CopyIntrinInjector : public StmtMutator { auto writer = dst.CopyOnWrite(); writer->shape = dst_shape; writer->strides = dst_strides; - writer->elem_offset = store_strides[loop_var_size]; + writer->elem_offsets = {store_strides[loop_var_size]}; } Buffer src = load->buffer; @@ -182,7 +182,7 @@ class CopyIntrinInjector : public StmtMutator { auto writer = src.CopyOnWrite(); writer->shape = src_shape; writer->strides = src_strides; - writer->elem_offset = src_elem_offset; + writer->elem_offsets = {src_elem_offset}; } *out = flower_copy_fromto_(src, dst, pad_before, pad_after, pad_value); if (!out->defined()) { diff --git a/src/tir/transforms/ir_utils.cc b/src/tir/transforms/ir_utils.cc index 700c9931bba0..37eeb0b8784f 100644 --- a/src/tir/transforms/ir_utils.cc +++ b/src/tir/transforms/ir_utils.cc @@ -170,8 +170,8 @@ class IRConvertSSA final : public StmtExprMutator { // new buffer, pushing it onto the scoped stack of existing // buffers. This will be popped when the new_buffer_var // redefinition is popped. - Buffer new_buf(new_buffer_var, buf->dtype, buf->shape, buf->strides, buf->elem_offset, - buf->name, buf->data_alignment, buf->offset_factor, buf->buffer_type, + Buffer new_buf(new_buffer_var, buf->dtype, buf->shape, buf->strides, buf->elem_offsets, + buf->name, buf->data_alignment, buf->offset_factors, buf->buffer_type, buf->axis_separators, buf->span); buffers.push_back(new_buf); return new_buf; diff --git a/src/tir/transforms/lower_match_buffer.cc b/src/tir/transforms/lower_match_buffer.cc index 5bde5cb90e2b..a6d228448cee 100644 --- a/src/tir/transforms/lower_match_buffer.cc +++ b/src/tir/transforms/lower_match_buffer.cc @@ -164,11 +164,21 @@ class MatchBufferLower : public StmtExprMutator { << " required_alignment=" << buffer->data_alignment << ", provided_alignment=" << source_buffer->data_alignment; } - if (is_zero(buffer->elem_offset)) { - ICHECK(is_zero(source_buffer->elem_offset)) - << "Trying to bind a Buffer with offset into one without offset " - << " required elem_offset=" << buffer->elem_offset - << ", provided elem_offset=" << source_buffer->elem_offset; + + ICHECK_EQ(buffer->elem_offsets.size(), source_buffer->elem_offsets.size()) + << "Trying to bind buffer with different physical dimension, requires " + << buffer->elem_offsets.size() << "-d buffer, but provided " + << source_buffer->elem_offsets.size() << "-d buffer"; + + for (size_t i = 0; i < buffer->elem_offsets.size(); i++) { + auto buffer_offset = buffer->elem_offsets[i]; + if (is_zero(buffer_offset)) { + auto source_offset = source_buffer->elem_offsets[i]; + ICHECK(is_zero(source_offset)) + << "Trying to bind a Buffer with offset into one without offset " + << " required elem_offset=" << buffer_offset + << ", provided elem_offset=" << source_offset; + } } // Step.2. Update @@ -186,16 +196,15 @@ class MatchBufferLower : public StmtExprMutator { } Array buffer_start_indices = source_buffer->ElemOffset(indices); - if (buffer_start_indices.size() == 1) { - Bind(buffer->elem_offset, buffer_start_indices[0], buffer->name + ".elem_offset"); - CHECK(analyzer_.CanProve(truncmod(buffer->elem_offset, buffer->offset_factor) == 0)) + ICHECK_EQ(buffer_start_indices.size(), buffer->elem_offsets.size()); + ICHECK_EQ(buffer_start_indices.size(), buffer->offset_factors.size()); + for (size_t i = 0; i < buffer_start_indices.size(); i++) { + auto offset = buffer->elem_offsets[i]; + auto factor = buffer->offset_factors[i]; + Bind(offset, buffer_start_indices[0], buffer->name + ".elem_offset"); + CHECK(analyzer_.CanProve(truncmod(offset, factor) == 0)) << "The source elem_offset " << buffer_start_indices[0] - << " does not satisfy the offset_factor " << buffer->offset_factor << "."; - } else { - // Non-zero elem_offset is ill-defined for non-flat memory. - // If needed in the future, will require `Array - // elem_offsets`, with one offset for each flattened index. - Bind(buffer->elem_offset, 0); + << " does not satisfy the offset_factor " << factor << "."; } } diff --git a/src/tir/transforms/lower_thread_allreduce.cc b/src/tir/transforms/lower_thread_allreduce.cc index 7e09943d0185..4419cc3e8e72 100644 --- a/src/tir/transforms/lower_thread_allreduce.cc +++ b/src/tir/transforms/lower_thread_allreduce.cc @@ -143,8 +143,8 @@ class ThreadAllreduceBuilder final : public StmtExprMutator { auto it = var_remap_.find(op->buffer->data.get()); if (it != var_remap_.end()) { Buffer remapped_buffer(it->second, op->buffer->dtype, op->buffer->shape, - op->buffer->strides, op->buffer->elem_offset, op->buffer->name, - op->buffer->data_alignment, op->buffer->offset_factor, + op->buffer->strides, op->buffer->elem_offsets, op->buffer->name, + op->buffer->data_alignment, op->buffer->offset_factors, op->buffer->buffer_type, op->buffer->axis_separators, op->buffer->span); buf_remap_[op->buffer.get()] = remapped_buffer; @@ -179,9 +179,9 @@ class ThreadAllreduceBuilder final : public StmtExprMutator { auto it = var_remap_.find(store->buffer->data.get()); if (it != var_remap_.end()) { Buffer remapped_buffer(it->second, store->buffer->dtype, store->buffer->shape, - store->buffer->strides, store->buffer->elem_offset, + store->buffer->strides, store->buffer->elem_offsets, store->buffer->name, store->buffer->data_alignment, - store->buffer->offset_factor, store->buffer->buffer_type, + store->buffer->offset_factors, store->buffer->buffer_type, store->buffer->axis_separators, store->buffer->span); buf_remap_[store->buffer.get()] = remapped_buffer; return BufferStore(remapped_buffer, store->value, store->indices, store->span); diff --git a/src/tir/transforms/simplify.cc b/src/tir/transforms/simplify.cc index 7d4fac8d7b2d..cf0aef009cf6 100644 --- a/src/tir/transforms/simplify.cc +++ b/src/tir/transforms/simplify.cc @@ -93,7 +93,7 @@ class StmtSimplifier : public IRMutatorWithAnalyzer { if (const BufferLoadNode* load = op->value.as()) { if (load->buffer->data.same_as(op->buffer->data) && ArrayDeepEqual(load->indices, op->indices) && - tir::ExprDeepEqual()(load->buffer->elem_offset, op->buffer->elem_offset) && + ArrayDeepEqual(load->buffer->elem_offsets, op->buffer->elem_offsets) && ArrayDeepEqual(load->buffer->shape, op->buffer->shape) && ArrayDeepEqual(load->buffer->strides, op->buffer->strides)) { return Evaluate(0); diff --git a/src/tir/transforms/storage_flatten.cc b/src/tir/transforms/storage_flatten.cc index 0d57f7928f47..9abc3bdea87b 100644 --- a/src/tir/transforms/storage_flatten.cc +++ b/src/tir/transforms/storage_flatten.cc @@ -1426,9 +1426,11 @@ class StorageFlattener : public StmtExprMutator { } } + size_t n_flattened_dims = op->buffer->axis_separators.size() + 1; e.buffer = Buffer(op->buffer->data, op->buffer->dtype, op->buffer->shape, op->buffer->strides, - PrimExpr(), op->buffer->name, align, 0, kDefault, - op->buffer->axis_separators, op->buffer->span); + Array(n_flattened_dims, PrimExpr()), op->buffer->name, align, + Array(n_flattened_dims, IntImm(op->buffer->DefaultIndexType(), 0)), + kDefault, op->buffer->axis_separators, op->buffer->span); e.flattened_buffer = e.buffer.GetFlattenedBuffer(); // TODO(Lunderberg): Move the handling of boolean into a diff --git a/src/tir/transforms/storage_rewrite.cc b/src/tir/transforms/storage_rewrite.cc index 0534f31c3423..dcfd452e3e69 100644 --- a/src/tir/transforms/storage_rewrite.cc +++ b/src/tir/transforms/storage_rewrite.cc @@ -430,9 +430,10 @@ class StoragePlanRewriter : public StmtExprMutator { return it->second; } - Buffer remapped = Buffer(new_backing_array, buf->dtype, buf->shape, buf->strides, - buf->elem_offset, new_backing_array->name_hint, buf->data_alignment, - buf->offset_factor, buf->buffer_type, buf->axis_separators, buf->span); + Buffer remapped = + Buffer(new_backing_array, buf->dtype, buf->shape, buf->strides, buf->elem_offsets, + new_backing_array->name_hint, buf->data_alignment, buf->offset_factors, + buf->buffer_type, buf->axis_separators, buf->span); buffer_remap_[key] = remapped; return remapped; } diff --git a/src/tir/usmp/transform/convert_pool_allocations_to_offsets.cc b/src/tir/usmp/transform/convert_pool_allocations_to_offsets.cc index b73534090ab5..a2ad02d5501d 100644 --- a/src/tir/usmp/transform/convert_pool_allocations_to_offsets.cc +++ b/src/tir/usmp/transform/convert_pool_allocations_to_offsets.cc @@ -356,8 +356,8 @@ Buffer PoolAllocationToOffsetConverter::GetRemappedBuffer(Buffer original) { auto it = allocate_var_to_let_var_.find(original->data); if (it != allocate_var_to_let_var_.end()) { remapped = Buffer((*it).second, original->dtype, original->shape, original->strides, - original->elem_offset, original->name, original->data_alignment, - original->offset_factor, original->buffer_type, original->axis_separators, + original->elem_offsets, original->name, original->data_alignment, + original->offset_factors, original->buffer_type, original->axis_separators, original->span); }