[FIX][TIRX] Repair typed-buffer CI regressions
diff --git a/include/tvm/tirx/expr.h b/include/tvm/tirx/expr.h index d4a63b6..96500d4 100644 --- a/include/tvm/tirx/expr.h +++ b/include/tvm/tirx/expr.h
@@ -553,7 +553,7 @@ static void RegisterReflection() { namespace refl = tvm::ffi::reflection; refl::ObjectDef<BufferLoadNode>() - .def_ro("buffer", &BufferLoadNode::buffer) + .def_ro("buffer", &BufferLoadNode::buffer, refl::AttachFieldFlag::SEqHashDefRecursive()) .def_ro("indices", &BufferLoadNode::indices) .def_ro("predicate", &BufferLoadNode::predicate); }
diff --git a/include/tvm/tirx/stmt.h b/include/tvm/tirx/stmt.h index 9d64ac2..d95c1af 100644 --- a/include/tvm/tirx/stmt.h +++ b/include/tvm/tirx/stmt.h
@@ -213,7 +213,7 @@ static void RegisterReflection() { namespace refl = tvm::ffi::reflection; refl::ObjectDef<BufferStoreNode>() - .def_ro("buffer", &BufferStoreNode::buffer) + .def_ro("buffer", &BufferStoreNode::buffer, refl::AttachFieldFlag::SEqHashDefRecursive()) .def_ro("value", &BufferStoreNode::value) .def_ro("indices", &BufferStoreNode::indices) .def_ro("predicate", &BufferStoreNode::predicate); @@ -783,7 +783,7 @@ static void RegisterReflection() { namespace refl = tvm::ffi::reflection; refl::ObjectDef<BufferRegionNode>() - .def_ro("buffer", &BufferRegionNode::buffer) + .def_ro("buffer", &BufferRegionNode::buffer, refl::AttachFieldFlag::SEqHashDefRecursive()) .def_ro("region", &BufferRegionNode::region); }
diff --git a/python/tvm/s_tir/dlight/analysis/common_analysis.py b/python/tvm/s_tir/dlight/analysis/common_analysis.py index 1706189..9c29d49 100644 --- a/python/tvm/s_tir/dlight/analysis/common_analysis.py +++ b/python/tvm/s_tir/dlight/analysis/common_analysis.py
@@ -158,7 +158,7 @@ ) vbuf_extent = int(self.shape[-1]) & ~(int(self.shape[-1]) - 1) - return min(vlp_extent, vbuf_extent, vbits // self.buf_region.buffer.dtype.dtype.bits) + return min(vlp_extent, vbuf_extent, vbits // self.buf_region.buffer.dtype.bits) def __str__(self) -> str: return f"BufferInfo({self.buf_region})"
diff --git a/python/tvm/s_tir/dlight/cpu/reduction.py b/python/tvm/s_tir/dlight/cpu/reduction.py index cf02a1f..bbce215 100644 --- a/python/tvm/s_tir/dlight/cpu/reduction.py +++ b/python/tvm/s_tir/dlight/cpu/reduction.py
@@ -85,9 +85,7 @@ # Infer dtype from the last block's write buffer. last_block_stmt = sch.get(block_infos[-1].block_rv) - dtype_bits = ( - last_block_stmt.writes[0].buffer.dtype.dtype.bits if last_block_stmt.writes else 32 - ) + dtype_bits = last_block_stmt.writes[0].buffer.dtype.bits if last_block_stmt.writes else 32 # Determine vector lanes from target VLEN. vlen_bits = llvm_get_vector_width(target)
diff --git a/python/tvm/tirx/buffer.py b/python/tvm/tirx/buffer.py index 6259265..ef9cee6 100644 --- a/python/tvm/tirx/buffer.py +++ b/python/tvm/tirx/buffer.py
@@ -570,8 +570,8 @@ elem_offset = Var(f"{name}_elem_offset", shape_ty) storage_scope = scope if data is not None: - if not isinstance(data, tvm.ir.Var) or not isinstance(data.ty, PointerType): - raise TypeError("Buffer data must be a Var with PointerType") + if not isinstance(data, tvm.ir.Expr) or not isinstance(data.ty, PointerType): + raise TypeError("Buffer data must be an Expr with PointerType") if not isinstance(data.ty.element_type, PrimType): raise TypeError("Buffer data must point to a primitive type") storage_scope = data.ty.storage_scope
diff --git a/python/tvm/tirx/op.py b/python/tvm/tirx/op.py index d3a6fcd..7812377 100644 --- a/python/tvm/tirx/op.py +++ b/python/tvm/tirx/op.py
@@ -1287,7 +1287,7 @@ call_args = [_pack_buffer(x) if is_buffer_var(x) else x for x in args] call_args.insert(0, tvm.tirx.StringImm(trace_action)) tracing_value = args[-1] - ret_ty = tracing_value.ty if isinstance(tracing_value, Expr) else tracing_value.ty.dtype + ret_ty = tracing_value.ty if isinstance(tracing_value, Expr) else tracing_value.dtype return tvm.ir.Call(Op.get("tirx.tvm_call_trace_packed"), call_args, ret_ty=ret_ty)
diff --git a/src/backend/metal/codegen/codegen_metal.cc b/src/backend/metal/codegen/codegen_metal.cc index c4f180f..78a0d74 100644 --- a/src/backend/metal/codegen/codegen_metal.cc +++ b/src/backend/metal/codegen/codegen_metal.cc
@@ -43,6 +43,25 @@ namespace tvm { namespace codegen { +namespace { + +Var GetSimdgroupBufferVar(const Expr& data) { + if (const auto* var = data.as<VarNode>()) { + return ffi::GetRef<Var>(var); + } + if (const auto* call = data.as<CallNode>(); + call && call->op.same_as(tirx::builtin::buffer_data()) && call->args.size() == 1) { + const auto* buffer = call->args[0].as<VarNode>(); + TVM_FFI_ICHECK(buffer && buffer->ty.as<BufferTypeNode>()) + << "Metal simdgroup data operands expect buffer_data to project a BufferVar"; + return ffi::GetRef<Var>(buffer); + } + TVM_FFI_THROW(InternalError) + << "Metal simdgroup data operands must be a Var or buffer_data(BufferVar), but got " << data; +} + +} // namespace + void CodeGenMetal::InitFuncState(const PrimFunc& f) { CodeGenC::InitFuncState(f); // analyze the data; @@ -388,7 +407,7 @@ if (op->op.same_as(make_filled_simdgroup_matrix_op)) { TVM_FFI_ICHECK_EQ(op->args.size(), 5); - Var var = op->args[0].as_or_throw<Var>(); + Var var = GetSimdgroupBufferVar(op->args[0]); // Get the data type of the simdgroup matrix auto it = simdgroup_dtype_.find(var.get()); TVM_FFI_ICHECK(it != simdgroup_dtype_.end()) @@ -401,25 +420,31 @@ << PrintExpr(op->args[2]) << ")"; } else if (op->op.same_as(simdgroup_load_op)) { TVM_FFI_ICHECK_EQ(op->args.size(), 7); + Var var = GetSimdgroupBufferVar(op->args[0]); f_check_simdgroup_shape(op->args[4].as_or_throw<PrimExpr>(), op->args[5].as_or_throw<PrimExpr>()); - os << "simdgroup_load(" << PrintExpr(op->args[0]) << "[" << PrintExpr(op->args[1]) << "], " + os << "simdgroup_load(" << PrintExpr(var) << "[" << PrintExpr(op->args[1]) << "], " << PrintExpr(op->args[2]) << ", " << PrintExpr(op->args[3]) << ", 0, " << PrintExpr(op->args[6]) << ")"; } else if (op->op.same_as(simdgroup_store_op)) { TVM_FFI_ICHECK_EQ(op->args.size(), 7); + Var var = GetSimdgroupBufferVar(op->args[0]); f_check_simdgroup_shape(op->args[4].as_or_throw<PrimExpr>(), op->args[5].as_or_throw<PrimExpr>()); - os << "simdgroup_store(" << PrintExpr(op->args[0]) << "[" << PrintExpr(op->args[1]) << "], " + os << "simdgroup_store(" << PrintExpr(var) << "[" << PrintExpr(op->args[1]) << "], " << PrintExpr(op->args[2]) << ", " << PrintExpr(op->args[3]) << ", 0, " << PrintExpr(op->args[6]) << ")"; } else if (op->op.same_as(simdgroup_multiply_accumulate_op)) { TVM_FFI_ICHECK_EQ(op->args.size(), 8); - os << "simdgroup_multiply_accumulate(" // - << PrintExpr(op->args[0]) << "[" << PrintExpr(op->args[1]) << "], " // - << PrintExpr(op->args[2]) << "[" << PrintExpr(op->args[3]) << "], " // - << PrintExpr(op->args[4]) << "[" << PrintExpr(op->args[5]) << "], " // - << PrintExpr(op->args[6]) << "[" << PrintExpr(op->args[7]) << "])"; + Var d = GetSimdgroupBufferVar(op->args[0]); + Var a = GetSimdgroupBufferVar(op->args[2]); + Var b = GetSimdgroupBufferVar(op->args[4]); + Var c = GetSimdgroupBufferVar(op->args[6]); + os << "simdgroup_multiply_accumulate(" // + << PrintExpr(d) << "[" << PrintExpr(op->args[1]) << "], " // + << PrintExpr(a) << "[" << PrintExpr(op->args[3]) << "], " // + << PrintExpr(b) << "[" << PrintExpr(op->args[5]) << "], " // + << PrintExpr(c) << "[" << PrintExpr(op->args[7]) << "])"; } else if (op->op.same_as(builtin::reinterpret())) { if (!op->ty.as<PrimTypeNode>() || !op->args[0]->ty.as<PrimTypeNode>()) { return CodeGenC::VisitExpr_(op, os);
diff --git a/src/backend/opencl/codegen/codegen_opencl.cc b/src/backend/opencl/codegen/codegen_opencl.cc index 8379c44..3ce33df 100644 --- a/src/backend/opencl/codegen/codegen_opencl.cc +++ b/src/backend/opencl/codegen/codegen_opencl.cc
@@ -37,6 +37,40 @@ namespace tvm { namespace codegen { +namespace { + +const VarNode* TryUnwrapTextureVar(const Expr& texture) { + if (const auto* var = texture.as<VarNode>()) { + return var; + } + if (const auto* call = texture.as<CallNode>(); call && call->op.same_as(builtin::buffer_data())) { + TVM_FFI_ICHECK_EQ(call->args.size(), 1U); + const auto* buffer = call->args[0].as<VarNode>(); + TVM_FFI_ICHECK(buffer && buffer->ty.as<BufferTypeNode>()) + << "buffer_data expects a Var with BufferType"; + return buffer; + } + return nullptr; +} + +struct TextureArgument { + const VarNode* var; + const PointerTypeNode* pointer_type; +}; + +TextureArgument UnwrapTextureArgument(const Expr& texture) { + const auto* var = TryUnwrapTextureVar(texture); + TVM_FFI_ICHECK(var) + << "Texture arguments must be a pointer Var or a buffer_data(BufferVar) projection"; + const auto* pointer_type = texture->ty.as<PointerTypeNode>(); + TVM_FFI_ICHECK(pointer_type) << "Texture arguments must have PointerType"; + TVM_FFI_ICHECK(runtime::IsTextureStorage(std::string(pointer_type->storage_scope))) + << "Texture intrinsics only support texture buffers"; + return {var, pointer_type}; +} + +} // namespace + class InferTextureAccess : public StmtExprVisitor { public: static constexpr const uint8_t kReadAccess = 1; @@ -44,6 +78,7 @@ InferTextureAccess() {} using StmtExprVisitor::VisitExpr_; + using StmtExprVisitor::VisitStmt_; std::unordered_map<const VarNode*, std::string> Infer(const Stmt& n) { StmtExprVisitor::VisitStmt(n); std::unordered_map<const VarNode*, std::string> storage_scope_qualifiers; @@ -58,17 +93,29 @@ } return storage_scope_qualifiers; } + void VisitStmt_(const DeclBufferNode* op) final { + if (const VarNode* source = TryUnwrapTextureVar(op->data)) { + auto it = buffer_data_map_.find(source); + buffer_data_map_[op->buffer.get()] = it == buffer_data_map_.end() ? source : it->second; + } + StmtExprVisitor::VisitStmt_(op); + } void VisitExpr_(const CallNode* op) final { if (op->op.same_as(builtin::texture2d_load())) { - var_access_map_[op->args[0].as<VarNode>()] |= kReadAccess; + const VarNode* texture = UnwrapTextureArgument(op->args[0]).var; + auto it = buffer_data_map_.find(texture); + var_access_map_[it == buffer_data_map_.end() ? texture : it->second] |= kReadAccess; } else if (op->op.same_as(builtin::texture2d_store())) { - var_access_map_[op->args[0].as<VarNode>()] |= kWriteAccess; + const VarNode* texture = UnwrapTextureArgument(op->args[0]).var; + auto it = buffer_data_map_.find(texture); + var_access_map_[it == buffer_data_map_.end() ? texture : it->second] |= kWriteAccess; } StmtExprVisitor::VisitExpr_(op); } private: std::unordered_map<const VarNode*, uint8_t> var_access_map_; + std::unordered_map<const VarNode*, const VarNode*> buffer_data_map_; }; CodeGenOpenCL::CodeGenOpenCL() { @@ -428,16 +475,13 @@ this->PrintExpr(load->indices[0], os); os << ')'; } else if (op->op.same_as(builtin::texture2d_store())) { - auto* ptr_type = op->args[0].as<VarNode>()->ty.as<PointerTypeNode>(); - TVM_FFI_ICHECK(ptr_type != nullptr) << "Texture Var's must be of PointerType"; - TVM_FFI_ICHECK(runtime::IsTextureStorage(std::string(ptr_type->storage_scope))) - << "builtin::texture2d_store() only supports storing to texture buffers"; + TextureArgument texture = UnwrapTextureArgument(op->args[0]); const int channel_size = op->args[4].as_or_throw<IntImm>()->value; TVM_FFI_ICHECK(channel_size == 64 || channel_size == 128) << "Unsupported Channel Size: " << channel_size; PrimType channel_type(runtime::GetChannelType(channel_size)); - PrimType buffer_type = ptr_type->element_type.as_or_throw<PrimType>(); + PrimType buffer_type = texture.pointer_type->element_type.as_or_throw<PrimType>(); std::stringstream ss; this->PrintExpr(op->args[5].as_or_throw<PrimExpr>(), ss); std::string value; @@ -449,7 +493,7 @@ } else { TVM_FFI_THROW(InternalError) << "Unsupported Channel Size: " << channel_size; } - this->PrintExpr(op->args[0], os); + os << this->GetVarID(texture.var); os << ", "; os << "(int4)("; this->PrintExpr(op->args[1].as_or_throw<PrimExpr>(), os); @@ -465,6 +509,7 @@ os << "(" << value << ")"; os << ")"; } else if (op->op.same_as(builtin::texture2d_load())) { + TextureArgument texture = UnwrapTextureArgument(op->args[0]); enable_compliant_texture_reads_ = true; std::stringstream ss; const int channel_size = op->args[4].as_or_throw<IntImm>()->value; @@ -482,7 +527,7 @@ } else { TVM_FFI_THROW(InternalError) << "Unsupported Channel Size: " << channel_size; } - this->PrintExpr(op->args[0], ss); + ss << this->GetVarID(texture.var); ss << ", "; ss << "image_sampler, "; ss << "((int4)(";
diff --git a/src/backend/trn/codegen/codegen_trn.cc b/src/backend/trn/codegen/codegen_trn.cc index fefccbf..fbad6be 100644 --- a/src/backend/trn/codegen/codegen_trn.cc +++ b/src/backend/trn/codegen/codegen_trn.cc
@@ -90,6 +90,7 @@ // clear previous generated state. this->InitFuncState(func); buffer_idmap_.clear(); + buffer_data_varmap_.clear(); data_buffer_idmap_.clear(); data_decl_buffer_map_.clear(); // skip the first underscore, so SSA variable starts from _1 @@ -114,6 +115,9 @@ LOG(FATAL) << "Trainium codegen currently only support buffer arguments"; }; std::string vid = AllocVarID(v.get()); + if (auto buffer = func->buffer_map.Get(v)) { + var_idmap_[buffer.value().get()] = vid; + } if (i >= static_cast<size_t>(num_inputs.value())) { this->stream << vid << ": nt.mutable_tensor, "; output_vids.push_back(vid); @@ -209,7 +213,7 @@ void CodeGenTrainium::VisitStmt_(const AllocBufferNode* op) { TVM_FFI_ICHECK(op->buffer.defined()); - std::string vid = AllocVarID(op->buffer.get()); + std::string vid = AllocVarID(op->buffer.get(), op->buffer.name() + "_ptr"); this->PrintIndent(); auto scope = op->buffer.scope(); @@ -607,7 +611,24 @@ if (op->buffer.scope() == "trn.psum" || op->buffer.scope() == "trn.sbuf") { return; } - const VarNode* data = op->buffer.get(); + const VarNode* data = op->data.as<VarNode>(); + if (const auto* call = op->data.as<CallNode>(); + call && call->op.same_as(builtin::buffer_data()) && call->args.size() == 1) { + data = call->args[0].as<VarNode>(); + } + TVM_FFI_ICHECK(data) << "Trainium codegen expects DeclBuffer data to be a buffer variable"; + if (data->ty.as<PointerTypeNode>()) { + buffer_idmap_[op->buffer] = GetVarID(data); + buffer_data_varmap_[op->buffer] = data; + return; + } + TVM_FFI_ICHECK(data->ty.as<BufferTypeNode>()); + BufferVar source_buffer(ffi::GetRef<Var>(data)); + auto source_it = buffer_data_varmap_.find(source_buffer); + TVM_FFI_ICHECK(source_it != buffer_data_varmap_.end()) + << "Trainium codegen expects the source buffer to be declared before its alias"; + data = source_it->second; + auto it = data_buffer_idmap_.find(data); if (it != data_buffer_idmap_.end()) { const BufferVar& prev_buffer = data_decl_buffer_map_.at(data); @@ -620,6 +641,7 @@ std::string data_vid = GetVarID(data); std::string buffer_vid = name_supply_->FreshName(data_vid + "_buffer"); buffer_idmap_[op->buffer] = buffer_vid; + buffer_data_varmap_[op->buffer] = data; data_buffer_idmap_[data] = buffer_vid; data_decl_buffer_map_[data] = op->buffer; PrintIndent();
diff --git a/src/backend/trn/codegen/codegen_trn.h b/src/backend/trn/codegen/codegen_trn.h index 63d6712..cb05ff9 100644 --- a/src/backend/trn/codegen/codegen_trn.h +++ b/src/backend/trn/codegen/codegen_trn.h
@@ -81,6 +81,8 @@ NKIInstructionCtx ctx_; std::unordered_map<std::string, std::string> opcode_map_; std::unordered_map<BufferVar, std::string, ffi::ObjectPtrHash, ffi::ObjectPtrEqual> buffer_idmap_; + std::unordered_map<BufferVar, const VarNode*, ffi::ObjectPtrHash, ffi::ObjectPtrEqual> + buffer_data_varmap_; std::unordered_map<const VarNode*, std::string> data_buffer_idmap_; std::unordered_map<const VarNode*, BufferVar> data_decl_buffer_map_; bool is_outermost_loop_ = true;
diff --git a/src/backend/vulkan/codegen/codegen_spirv.cc b/src/backend/vulkan/codegen/codegen_spirv.cc index b35137e..69958d7 100644 --- a/src/backend/vulkan/codegen/codegen_spirv.cc +++ b/src/backend/vulkan/codegen/codegen_spirv.cc
@@ -57,6 +57,17 @@ return node; } +const VarNode* AsBufferVarNode(const Expr& expr) { + if (const auto* var = expr.as<VarNode>()) { + return var; + } + if (const auto* call = expr.as<CallNode>(); + call && call->op.same_as(builtin::buffer_data()) && call->args.size() == 1) { + return call->args[0].as<VarNode>(); + } + return nullptr; +} + } // namespace CodeGenSPIRV::CodeGenSPIRV(Target target) : spirv_support_(target) {} @@ -434,7 +445,7 @@ if (op->op.same_as(tvm_fill_fragment_op)) { TVM_FFI_ICHECK_EQ(op->args.size(), 6U); - const VarNode* buffer_node = op->args[0].as<VarNode>(); + const VarNode* buffer_node = AsBufferVarNode(op->args[0]); TVM_FFI_ICHECK(buffer_node && fragment_info_.count(buffer_node)); PrimType ele_dtype = GetElementDataType(buffer_node); TVM_FFI_ICHECK(ele_dtype.MatchesCode(DLDataTypeCode::kDLFloat)) @@ -454,7 +465,7 @@ } else if (op->op.same_as(tvm_load_matrix_sync_op)) { TVM_FFI_ICHECK_EQ(op->args.size(), 8U); - const VarNode* buffer_node = op->args[0].as<VarNode>(); + const VarNode* buffer_node = AsBufferVarNode(op->args[0]); TVM_FFI_ICHECK(buffer_node && fragment_info_.count(buffer_node)); spirv::SType& fragment_type = fragment_info_[buffer_node].stype; PrimExpr dst_index = op->args[4].as_or_throw<PrimExpr>(); @@ -476,10 +487,12 @@ builder_->MakeInst(spv::OpStore, dst_ptr, loaded, spv::MemoryAccessMaskNone); return spirv::Value(); } else if (op->op.same_as(tvm_mma_sync_op)) { - const VarNode* buffer_d = op->args[0].as<VarNode>(); - const VarNode* buffer_a = op->args[2].as<VarNode>(); - const VarNode* buffer_b = op->args[4].as<VarNode>(); - const VarNode* buffer_c = op->args[6].as<VarNode>(); + const VarNode* buffer_d = AsBufferVarNode(op->args[0]); + const VarNode* buffer_a = AsBufferVarNode(op->args[2]); + const VarNode* buffer_b = AsBufferVarNode(op->args[4]); + const VarNode* buffer_c = AsBufferVarNode(op->args[6]); + TVM_FFI_ICHECK(buffer_d && buffer_a && buffer_b && buffer_c) + << "Cooperative matrix operands must be buffer variables"; PrimExpr index_d = op->args[1].as_or_throw<PrimExpr>(); PrimExpr index_a = op->args[3].as_or_throw<PrimExpr>(); PrimExpr index_b = op->args[5].as_or_throw<PrimExpr>(); @@ -514,7 +527,8 @@ return spirv::Value(); } else if (op->op.same_as(tvm_store_matrix_sync_op)) { TVM_FFI_ICHECK_EQ(op->args.size(), 8U); - const VarNode* buffer_node = op->args[0].as<VarNode>(); + const VarNode* buffer_node = AsBufferVarNode(op->args[0]); + TVM_FFI_ICHECK(buffer_node && fragment_info_.count(buffer_node)); PrimExpr index = op->args[4].as_or_throw<PrimExpr>(); int stride = static_cast<int>(AsIntImmNode(op->args[6])->value); auto type_int = builder_->GetSType(PrimType::Int(32)); @@ -898,7 +912,35 @@ } void CodeGenSPIRV::VisitStmt_(const DeclBufferNode* op) { - // DeclBuffer is a flat statement with no body — nothing to emit. + const VarNode* buffer_var = op->buffer.get(); + TVM_FFI_ICHECK(!var_map_.count(buffer_var)) + << "Buffer variable " << op->buffer.name() << " is already defined"; + TVM_FFI_ICHECK(!storage_info_.count(buffer_var)) + << "Storage metadata for buffer variable " << op->buffer.name() << " is already defined"; + + spirv::Value data = MakeValue(op->data); + + PrimType declared_storage_type = op->buffer->dtype; + if (declared_storage_type == PrimType::Bool()) { + declared_storage_type = boolean_storage_type_.WithLanes(declared_storage_type.lanes()); + } + + const VarNode* source = AsBufferVarNode(op->data); + + StorageInfo info; + if (source) { + auto it = storage_info_.find(source); + if (it != storage_info_.end()) { + info = it->second; + info.name_hint = op->buffer.name(); + } + } + if (!info.element_type_known) { + info.SetContentType(declared_storage_type, op->buffer.name()); + } + + var_map_[buffer_var] = data; + storage_info_[buffer_var] = std::move(info); } void CodeGenSPIRV::VisitStmt_(const AttrStmtNode* op) {
diff --git a/src/relax/script/printer/expr.cc b/src/relax/script/printer/expr.cc index 3a63325..34ce7b0 100644 --- a/src/relax/script/printer/expr.cc +++ b/src/relax/script/printer/expr.cc
@@ -158,7 +158,8 @@ std::string ReprPrintVar(const ffi::ObjectRef& obj, const PrinterConfig& cfg) { Var var = obj.as_or_throw<Var>(); - if (var->ty.as<PrimTypeNode>() || var->ty.as<PointerTypeNode>()) { + if (var->ty.as<PrimTypeNode>() || var->ty.as<PointerTypeNode>() || + var->ty.as<tirx::BufferTypeNode>()) { return ReprPrintTIR(obj, cfg); } return ReprPrintRelax(obj, cfg);
diff --git a/src/relax/transform/fuse_tir.cc b/src/relax/transform/fuse_tir.cc index 28f7178..a714514 100644 --- a/src/relax/transform/fuse_tir.cc +++ b/src/relax/transform/fuse_tir.cc
@@ -497,7 +497,7 @@ // structurally equal to the `new_buf` passed auto ValidateBufferCompatibility = [this](tirx::BufferVar new_buf, Expr expr) { if (auto it = relax_to_tir_var_map_.find(expr); it != relax_to_tir_var_map_.end()) { - TVM_FFI_ICHECK(ffi::StructuralEqual()((*it).second, new_buf)) + TVM_FFI_ICHECK(ffi::StructuralEqual()((*it).second.type(), new_buf.type())) << "Inconsistent buffers " << (*it).second << " and " << new_buf << " mapped to the same relax var: " << expr; }
diff --git a/src/s_tir/analysis/sblock_access_region_detector.cc b/src/s_tir/analysis/sblock_access_region_detector.cc index 4eb3f32..3feb2c4 100644 --- a/src/s_tir/analysis/sblock_access_region_detector.cc +++ b/src/s_tir/analysis/sblock_access_region_detector.cc
@@ -122,6 +122,7 @@ void VisitStmt_(const ForNode* op) override; void VisitStmt_(const IfThenElseNode* op) override; void VisitStmt_(const SBlockRealizeNode* op) override; + void VisitStmt_(const DeclBufferNode* op) override; void VisitStmt_(const BufferStoreNode* op) override; void VisitStmt_(const BindNode* op) override; void VisitExpr_(const BufferLoadNode* op) override; @@ -195,6 +196,12 @@ } } +void BlockReadWriteDetector::VisitStmt_(const DeclBufferNode* op) { + // A DeclBuffer data expression defines the alias source. It is not an + // opaque buffer access by the containing block. + VisitBufferDef(op->buffer, /*alloc_data=*/false); +} + void BlockReadWriteDetector::VisitStmt_(const BindNode* op) { if (auto value = op->value.as<PrimExpr>()) { let_bindings_[op->var.get()] = value.value();
diff --git a/src/s_tir/transform/lower_thread_allreduce.cc b/src/s_tir/transform/lower_thread_allreduce.cc index 1b28872..027c124d 100644 --- a/src/s_tir/transform/lower_thread_allreduce.cc +++ b/src/s_tir/transform/lower_thread_allreduce.cc
@@ -135,7 +135,8 @@ } ffi::Optional<BufferVar> GetRemappedBuffer(const BufferVar& buf) { - if (auto it = var_remap_.find(buf.get()); it != var_remap_.end()) { + Var root = buffer_aliases_.Get(buf.var()).value_or(buf.var()); + if (auto it = var_remap_.find(root.get()); it != var_remap_.end()) { return BufferVar(it->second); } @@ -144,15 +145,15 @@ Stmt VisitStmt_(const DeclBufferNode* op) final { RegisterBufferAlias(op->buffer, op->data); - auto node = StmtExprMutator::VisitStmt_(op).as_or_throw<DeclBuffer>(); - if (auto buf = GetRemappedBuffer(node->buffer)) { - node.CopyOnWrite()->buffer = buf.value(); - } - return node; + // Remap declarations only after the complete traversal has populated the + // physical-root maps. Eagerly replacing an alias declared after its + // allreduce would retain the old source pointer on the new buffer. + return StmtExprMutator::VisitStmt_(op); } Expr VisitExpr_(const BufferLoadNode* op) final { - if (auto it = load_remap_.find(op->buffer.get()); it != load_remap_.end()) { + const VarNode* allocation = GetAllocationKey(op->buffer.get()); + if (auto it = load_remap_.find(allocation); it != load_remap_.end()) { for (const auto& index : op->indices) { TVM_FFI_ICHECK(is_zero(index)); } @@ -169,9 +170,19 @@ } Stmt VisitStmt_(const BufferStoreNode* op) final { + const VarNode* allocation = GetAllocationKey(op->buffer.get()); BufferStore store = StmtExprMutator::VisitStmt_(op).as_or_throw<BufferStore>(); - if (auto opt = GetRemappedBuffer(store->buffer)) { + if (auto it = load_remap_.find(allocation); it != load_remap_.end()) { + const auto* replacement = it->second.as<BufferLoadNode>(); + TVM_FFI_ICHECK(replacement); + for (const auto& index : store->indices) { + TVM_FFI_ICHECK(is_zero(index)); + } + auto* writer = store.CopyOnWrite(); + writer->buffer = replacement->buffer; + writer->indices = replacement->indices; + } else if (auto opt = GetRemappedBuffer(store->buffer)) { store.CopyOnWrite()->buffer = opt.value(); } return store; @@ -415,14 +426,14 @@ // Write back allreduce results and update existing allocations. for (size_t i = 0; i < size; ++i) { - TVM_FFI_ICHECK(!load_remap_.count(buffers[i].get())); + const VarNode* alloc_key = GetAllocationKey(buffers[i].get()); + TVM_FFI_ICHECK(!load_remap_.count(alloc_key)); BufferVar buf = reduce_results[i].as_or_throw<BufferLoad>()->buffer; TVM_FFI_ICHECK_EQ(reduce_results[i].ty(), dtypes[i]); - load_remap_[buffers[i].get()] = reduce_results[i]; + load_remap_[alloc_key] = reduce_results[i]; // The AllocBuffer doesn't need to be emitted here since alloc_remap_ // will cause the existing allocation to be rewritten in VisitStmt_(AllocBufferNode*). - const VarNode* alloc_key = GetAllocationKey(buffers[i].get()); alloc_remap_[alloc_key] = buf; var_remap_[alloc_key] = buf.var(); var_remap_[buffers[i].get()] = buf.var(); @@ -450,13 +461,13 @@ seq.emplace_back(MakeBufAllreduce(combiner, dtypes, shared_bufs, reduce_index, group_index, reduce_extent, group_extent, contiguous_reduce_extent)); for (size_t idx = 0; idx < size; ++idx) { - TVM_FFI_ICHECK(!load_remap_.count(buffers[idx].get())); + const VarNode* alloc_key = GetAllocationKey(buffers[idx].get()); + TVM_FFI_ICHECK(!load_remap_.count(alloc_key)); PrimExpr pred = MakeConst(PrimType::Bool(static_cast<int16_t>(dtypes[idx].lanes())), true); BufferLoad load(shared_bufs[idx], {BufIndex(IntImm(reduce_index.ty(), 0), group_index, reduce_extent)}); TVM_FFI_ICHECK_EQ(load.ty(), dtypes[idx]); - load_remap_[buffers[idx].get()] = load; - const VarNode* alloc_key = GetAllocationKey(buffers[idx].get()); + load_remap_[alloc_key] = load; alloc_remap_[alloc_key] = shared_bufs[idx]; var_remap_[alloc_key] = shared_bufs[idx].var(); var_remap_[buffers[idx].get()] = shared_bufs[idx].var(); @@ -936,7 +947,8 @@ private: ffi::Optional<BufferVar> GetRemappedBuffer(const BufferVar& buf) { - if (auto it = var_remap_.find(buf.get()); it != var_remap_.end()) { + Var root = buffer_aliases_.Get(buf.var()).value_or(buf.var()); + if (auto it = var_remap_.find(root.get()); it != var_remap_.end()) { return BufferVar(it->second); } return std::nullopt;
diff --git a/src/s_tir/transform/storage_access.cc b/src/s_tir/transform/storage_access.cc index 4db74ee..897ccfa 100644 --- a/src/s_tir/transform/storage_access.cc +++ b/src/s_tir/transform/storage_access.cc
@@ -51,7 +51,7 @@ } // namespace void StorageAccessVisitor::VisitExpr_(const BufferLoadNode* op) { - Var buf = op->buffer.var(); + Var buf = ResolveBuffer(op->buffer.var()); StorageScope scope = StorageScope::Create(op->buffer.scope()); if (Enabled(buf.get(), scope)) { TVM_FFI_ICHECK(allow_append_) << op << " " << scope.to_string(); @@ -75,7 +75,7 @@ TVM_FFI_ICHECK_EQ(curr_stmt_.access.size(), 0U); curr_stmt_.stmt = op; - Var buf = op->buffer.var(); + Var buf = ResolveBuffer(op->buffer.var()); StorageScope scope = StorageScope::Create(op->buffer.scope()); if (Enabled(buf.get(), scope)) { AccessEntry e; @@ -98,6 +98,13 @@ allow_append_ = false; } +void StorageAccessVisitor::VisitStmt_(const DeclBufferNode* op) { + if (auto source = GetBufferDataVar(op->data)) { + buffer_aliases_.insert_or_assign(op->buffer.get(), ResolveBuffer(source.value())); + } + StmtExprVisitor::VisitStmt_(op); +} + void StorageAccessVisitor::VisitStmt_(const EvaluateNode* op) { allow_append_ = true; TVM_FFI_ICHECK_EQ(curr_stmt_.access.size(), 0U); @@ -129,7 +136,7 @@ auto buffer = GetBufferDataVar(op->node); TVM_FFI_ICHECK(buffer.has_value()) << "Expected a buffer data expression for double-buffer writes, but received " << op->node; - double_buffer_write_ = buffer.value().get(); + double_buffer_write_ = ResolveBuffer(buffer.value()).get(); scope_.push_back(std::vector<StmtEntry>()); StmtExprVisitor::VisitStmt_(op); StmtEntry s; @@ -275,18 +282,18 @@ StmtExprVisitor::VisitExpr_(op); return; } - const VarNode* buffer = buffer_var.value().get(); + Var buffer = ResolveBuffer(buffer_var.value()); PrimExpr offset = op->args[2].as_or_throw<PrimExpr>(); PrimExpr extent = op->args[3].as_or_throw<PrimExpr>(); const IntImmNode* flag = op->args[4].as<IntImmNode>(); StorageScope scope = GetScope(buffer_var.value()); // The buffer scope. - if (Enabled(buffer, scope)) { + if (Enabled(buffer.get(), scope)) { TVM_FFI_ICHECK(allow_append_); AccessEntry e; e.threads = env_threads(); e.dtype = dtype; - e.buffer = ffi::GetRef<Var>(buffer); + e.buffer = buffer; e.touched = {arith::IntSet::FromRange(Range::FromMinExtent(offset, extent))}; e.scope = scope; if (flag->value & 1) { @@ -325,5 +332,10 @@ return StorageScope(); // global by default } +Var StorageAccessVisitor::ResolveBuffer(Var buffer_var) const { + auto it = buffer_aliases_.find(buffer_var.get()); + return it == buffer_aliases_.end() ? buffer_var : it->second; +} + } // namespace s_tir } // namespace tvm
diff --git a/src/s_tir/transform/storage_access.h b/src/s_tir/transform/storage_access.h index 8442947..d7c8de3 100644 --- a/src/s_tir/transform/storage_access.h +++ b/src/s_tir/transform/storage_access.h
@@ -84,6 +84,7 @@ // override visitor pattern void VisitExpr_(const BufferLoadNode* op) final; void VisitStmt_(const BufferStoreNode* op) final; + void VisitStmt_(const DeclBufferNode* op) final; void VisitStmt_(const EvaluateNode* op) final; void VisitStmt_(const BindNode* op) final; void VisitStmt_(const AttrStmtNode* op) final; @@ -124,6 +125,8 @@ * \return The scope of the final buffer array. */ StorageScope GetScope(Var buffer_var) const; + /*! \brief Resolve a logical buffer view to its physical storage root. */ + Var ResolveBuffer(Var buffer_var) const; // access scope std::vector<std::vector<StmtEntry>> scope_; @@ -140,6 +143,8 @@ StmtEntry curr_stmt_; // The involving threads ffi::Array<IterVar> env_threads_; + // Physical storage root for each declared logical buffer view. + std::unordered_map<const VarNode*, Var> buffer_aliases_; }; } // namespace s_tir } // namespace tvm
diff --git a/src/s_tir/transform/tensorcore_infer_fragment.cc b/src/s_tir/transform/tensorcore_infer_fragment.cc index a97d5a1..402848c 100644 --- a/src/s_tir/transform/tensorcore_infer_fragment.cc +++ b/src/s_tir/transform/tensorcore_infer_fragment.cc
@@ -41,6 +41,17 @@ namespace s_tir { using namespace tvm::tirx; +const VarNode* GetBufferVarFromData(const Expr& data) { + if (const auto* var = data.as<VarNode>()) { + return var; + } + if (const auto* call = data.as<CallNode>(); + call && call->op.same_as(builtin::buffer_data()) && call->args.size() == 1) { + return call->args[0].as<VarNode>(); + } + return nullptr; +} + // Get fragment information from tensor intrinsics class FragmentGetter : public StmtExprVisitor { public: @@ -53,7 +64,7 @@ if (op->op.same_as(tvm_load_matrix_sync_op) || op->op.same_as(tvm_store_matrix_sync_op)) { // Get shape and layout information from load and store intrinsic TVM_FFI_ICHECK_EQ(op->args.size(), 8U); - const VarNode* buffer_var = op->args[0].as<VarNode>(); + const VarNode* buffer_var = GetBufferVarFromData(op->args[0]); TVM_FFI_ICHECK(buffer_var); // Get shape const IntImmNode* m = op->args[1].as<IntImmNode>(); @@ -88,7 +99,7 @@ } else if (op->op.same_as(tvm_fill_fragment_op)) { // Get shape information from fill intrinsic TVM_FFI_ICHECK_EQ(op->args.size(), 6U); - const VarNode* buffer_var = op->args[0].as<VarNode>(); + const VarNode* buffer_var = GetBufferVarFromData(op->args[0]); TVM_FFI_ICHECK(buffer_var); // Get shape const IntImmNode* m = op->args[1].as<IntImmNode>(); @@ -143,10 +154,10 @@ static const Op& tvm_bmma_sync_op = Op::Get("tirx.tvm_bmma_sync"); if (op->op.same_as(tvm_mma_sync_op) || op->op.same_as(tvm_bmma_sync_op)) { TVM_FFI_ICHECK_EQ(op->args.size(), 8U); - const VarNode* buffer_var_d = op->args[0].as<VarNode>(); - const VarNode* buffer_var_a = op->args[2].as<VarNode>(); - const VarNode* buffer_var_b = op->args[4].as<VarNode>(); - const VarNode* buffer_var_c = op->args[6].as<VarNode>(); + const VarNode* buffer_var_d = GetBufferVarFromData(op->args[0]); + const VarNode* buffer_var_a = GetBufferVarFromData(op->args[2]); + const VarNode* buffer_var_b = GetBufferVarFromData(op->args[4]); + const VarNode* buffer_var_c = GetBufferVarFromData(op->args[6]); TVM_FFI_ICHECK(buffer_var_d); TVM_FFI_ICHECK(buffer_var_a); TVM_FFI_ICHECK(buffer_var_b);
diff --git a/src/target/llvm/codegen_llvm.cc b/src/target/llvm/codegen_llvm.cc index cc4e947..f8fd8d1 100644 --- a/src/target/llvm/codegen_llvm.cc +++ b/src/target/llvm/codegen_llvm.cc
@@ -241,6 +241,7 @@ void CodeGenLLVM::InitFuncState() { var_map_.clear(); + buffer_physical_root_.clear(); alias_var_set_.clear(); alloc_storage_info_.clear(); volatile_buf_.clear(); @@ -1723,6 +1724,11 @@ return bytes != bytes_scalar * dtype.lanes(); } +const VarNode* CodeGenLLVM::GetBufferPhysicalRoot(const VarNode* buffer) const { + auto it = buffer_physical_root_.find(buffer); + return it == buffer_physical_root_.end() ? buffer : it->second; +} + void CodeGenLLVM::BufferAccessHelper( BufferVar buffer, ffi::Array<PrimExpr> indices, ffi::Optional<PrimExpr> predicate, PrimType value_dtype, @@ -1756,7 +1762,8 @@ PrimExpr last_index_origin = last_index; PrimType buffer_element_dtype_origin = buffer_element_dtype; - bool is_volatile = volatile_buf_.count(buffer.get()); + const VarNode* physical_root = GetBufferPhysicalRoot(buffer.get()); + bool is_volatile = volatile_buf_.count(physical_root); // If the buffer index is a contiguous ramp node, we only need to // access the first element, then cast to the value type. @@ -1784,7 +1791,7 @@ // element being accessed may require more alignment than the // underlying data type. int native_bits; - GetAlignment(value_dtype, buffer.get(), last_index, &alignment, &native_bits); + GetAlignment(value_dtype, physical_root, last_index, &alignment, &native_bits); } else { // Otherwise, alignment is based on the return value's scalar // type. @@ -1828,7 +1835,7 @@ value_dtype.WithLanes(value_dtype.lanes() / last_index_lanes)); auto instruction = make_instruction(buffer_ptr, subelement_i, predicate_value, alignment, is_volatile); - AddAliasInfo(instruction, buffer.get(), last_index_origin, buffer_element_dtype_origin); + AddAliasInfo(instruction, physical_root, last_index_origin, buffer_element_dtype_origin); } } @@ -2229,6 +2236,15 @@ } llvm::Value* value = MakeValue(op->data); + const VarNode* source = op->data.as<VarNode>(); + if (const auto* call = op->data.as<CallNode>(); + call && call->op.same_as(builtin::buffer_data()) && call->args.size() == 1) { + source = call->args[0].as<VarNode>(); + } + if (source) { + buffer_physical_root_[buffer] = GetBufferPhysicalRoot(source); + } + llvm::Type* expected_type = GetLLVMType(op->buffer.DataPointerType()); if (value->getType() != expected_type) { value->setName((op->buffer.name() + "_source_ptr").c_str()); @@ -2237,9 +2253,11 @@ AddDebugInformation(value, op->buffer.var()); var_map_[buffer] = value; - if (alloc_storage_info_.count(buffer) && alloc_storage_info_[buffer].alignment > 1) { + const VarNode* physical_root = GetBufferPhysicalRoot(buffer); + if (alloc_storage_info_.count(physical_root) && + alloc_storage_info_[physical_root].alignment > 1) { builder_->CreateAlignmentAssumption(*data_layout_, GetVarValue(buffer), - alloc_storage_info_[buffer].alignment); + alloc_storage_info_[physical_root].alignment); } }
diff --git a/src/target/llvm/codegen_llvm.h b/src/target/llvm/codegen_llvm.h index 41b237a..8a7e0bd 100644 --- a/src/target/llvm/codegen_llvm.h +++ b/src/target/llvm/codegen_llvm.h
@@ -363,6 +363,7 @@ std::function<llvm::Instruction*(TypedPointer buffer_ptr, int subelement_i, llvm::Value* predicate, int alignment, bool is_volatile)> make_instruction); + const VarNode* GetBufferPhysicalRoot(const VarNode* buffer) const; // Initialize target virtual void InitTarget(); // Add module startup function if needed. @@ -548,6 +549,8 @@ std::unordered_map<const VarNode*, StorageInfo> alloc_storage_info_; // The definition of local variable. std::unordered_map<const VarNode*, llvm::Value*> var_map_; + // Canonical physical storage identity for DeclBuffer aliases. + std::unordered_map<const VarNode*, const VarNode*> buffer_physical_root_; // global strings std::unordered_map<std::string, llvm::Constant*> str_map_;
diff --git a/src/target/source/codegen_c.cc b/src/target/source/codegen_c.cc index bec1342..423f239 100644 --- a/src/target/source/codegen_c.cc +++ b/src/target/source/codegen_c.cc
@@ -829,6 +829,9 @@ } else if (op->op.same_as(builtin::reinterpret())) { if (const auto* pointer_type = op->ty.as<PointerTypeNode>()) { os << "(("; + if (IsScopePartOfType()) { + PrintStorageScope(pointer_type->storage_scope, os); + } if (const auto* element_type = pointer_type->element_type.as<PrimTypeNode>()) { this->PrintType(ffi::GetRef<PrimType>(element_type), os); } else { @@ -918,11 +921,26 @@ if (source && var_idmap_.count(source)) { TVM_FFI_ICHECK(!var_idmap_.count(op->buffer.get())); var_idmap_[op->buffer.get()] = GetVarID(source); - RegisterHandleType(op->buffer.get(), op->buffer->dtype); + if (auto it = alloc_storage_scope_.find(source); it != alloc_storage_scope_.end()) { + alloc_storage_scope_[op->buffer.get()] = it->second; + } else { + alloc_storage_scope_[op->buffer.get()] = op->buffer.scope(); + } + if (IsVolatile(source)) { + MarkVolatile(op->buffer.get()); + } + auto it = handle_data_type_.find(source); + RegisterHandleType(op->buffer.get(), + it == handle_data_type_.end() ? op->buffer->dtype : it->second); return; } + std::string scope = op->buffer.scope(); + alloc_storage_scope_[op->buffer.get()] = scope; this->PrintIndent(); + if (IsScopePartOfType()) { + PrintStorageScope(scope, stream); + } PrintType(op->buffer.DataPointerType(), stream); stream << ' ' << AllocVarID(op->buffer.get()) << " = "; PrintExpr(Call(op->buffer.DataPointerType(), builtin::reinterpret(), {op->data}), stream);
diff --git a/src/target/source/codegen_c_host.h b/src/target/source/codegen_c_host.h index e9b89e6..94fd3a5 100644 --- a/src/target/source/codegen_c_host.h +++ b/src/target/source/codegen_c_host.h
@@ -75,6 +75,11 @@ const Type& ret_type) override; ffi::Array<ffi::String> GetFunctionNames() { return function_names_; } + protected: + // Plain C has no address-space qualifiers. Buffer storage scopes still + // control allocation lowering, but are not part of emitted pointer types. + bool IsScopePartOfType() const final { return false; } + private: std::string module_name_; /* \brief mapping global packed func to the unique name */
diff --git a/src/te/operation/create_primfunc.cc b/src/te/operation/create_primfunc.cc index 410dc4c..90c8922 100644 --- a/src/te/operation/create_primfunc.cc +++ b/src/te/operation/create_primfunc.cc
@@ -655,6 +655,7 @@ input_buffer_map[placeholder.get()] = output_buffer; info->root_alloc.push_back(output_buffer); } + var_map[placeholder.get()] = output_buffer.var(); info->tensor2buffers[output_tensor] = output_buffer; }
diff --git a/src/tirx/ir/data_type_rewriter.cc b/src/tirx/ir/data_type_rewriter.cc index 04abe0d..a8cb4f6 100644 --- a/src/tirx/ir/data_type_rewriter.cc +++ b/src/tirx/ir/data_type_rewriter.cc
@@ -164,6 +164,9 @@ } Expr DataTypeLegalizer::VisitExpr_(const VarNode* op) { + if (op->ty.as<BufferTypeNode>()) { + return VisitBufferUse(GetBufferVar(op)).var(); + } if (auto it = var_remap_.find(op); it != var_remap_.end()) { return it->second; } @@ -392,8 +395,8 @@ if (obj == nullptr) { return obj; } - if (obj.as<BufferTypeNode>()) { - BufferVar buffer = obj.as_or_throw<BufferVar>(); + if (auto var = obj.as<Var>(); var && var.value()->ty.as<BufferTypeNode>()) { + BufferVar buffer(var.value()); if (BufferVar new_buffer = VisitBufferUse(buffer); !new_buffer.same_as(buffer)) { return new_buffer; }
diff --git a/src/tirx/ir/specialize.cc b/src/tirx/ir/specialize.cc index 804eec2..687b37d 100644 --- a/src/tirx/ir/specialize.cc +++ b/src/tirx/ir/specialize.cc
@@ -189,6 +189,16 @@ private: BufferVar MutateBuffer(const BufferVar& buffer) { + ffi::Optional<ffi::String> specialized_storage_scope; + if (auto it = var_map_.find(buffer.var()); it != var_map_.end()) { + if (const auto* new_var = it->second.as<VarNode>()) { + if (new_var->ty.as<BufferTypeNode>()) { + BufferVar replacement(ffi::GetRef<Var>(new_var)); + specialized_storage_scope = replacement->storage_scope; + } + } + } + ffi::Array<PrimExpr> shape = buffer->shape.Map([this](const PrimExpr& e) { return VisitPrimExpr(e); }); ffi::Array<PrimExpr> strides = @@ -220,8 +230,10 @@ } } + bool storage_scope_changed = specialized_storage_scope.has_value() && + specialized_storage_scope.value() != buffer->storage_scope; if (buffer->elem_offset.same_as(elem_offset) && buffer->shape.same_as(shape) && - buffer->strides.same_as(strides) && !layout_changed) { + buffer->strides.same_as(strides) && !layout_changed && !storage_scope_changed) { return buffer; } else { auto n = CopyBufferType(buffer); @@ -231,6 +243,9 @@ if (layout_changed) { n->layout = std::move(layout); } + if (storage_scope_changed) { + n->storage_scope = specialized_storage_scope.value(); + } return RebuildBufferVar(buffer, std::move(n)); } }
diff --git a/src/tirx/ir/stmt.cc b/src/tirx/ir/stmt.cc index 14ca7f7..817da4e 100644 --- a/src/tirx/ir/stmt.cc +++ b/src/tirx/ir/stmt.cc
@@ -382,6 +382,9 @@ // Evaluate Evaluate::Evaluate(Expr value, Span span) { TVM_FFI_ICHECK(value.defined()); + TVM_FFI_ICHECK(!(value->IsInstance<VarNode>() && value->ty.as<BufferTypeNode>())) + << "A buffer variable cannot be used as a scalar Evaluate value; " + << "use buffer.data to evaluate its physical pointer"; ffi::ObjectPtr<EvaluateNode> node = ffi::make_object<EvaluateNode>(); node->value = std::move(value);
diff --git a/src/tirx/script/printer/expr.cc b/src/tirx/script/printer/expr.cc index 87915fa..9d79955 100644 --- a/src/tirx/script/printer/expr.cc +++ b/src/tirx/script/printer/expr.cc
@@ -46,7 +46,7 @@ } else { ExprDoc element_type = LiteralDoc::DataType(prim_type->dtype, type_p->Attr("element_type")->Attr("dtype")); - if (ptr_type->storage_scope == "global") { + if (ptr_type->storage_scope.empty()) { rhs = rhs->Call({element_type}, kwargs_keys, kwargs_values); } else { rhs = rhs->Call({element_type,
diff --git a/src/tirx/script/printer/stmt.cc b/src/tirx/script/printer/stmt.cc index 5b6c30e..9aa9822 100644 --- a/src/tirx/script/printer/stmt.cc +++ b/src/tirx/script/printer/stmt.cc
@@ -176,7 +176,8 @@ TVM_STATIC_IR_FUNCTOR(IRDocsifier, vtable) .set_dispatch<tirx::Evaluate>("", [](tirx::Evaluate eval, AccessPath p, IRDocsifier d) -> Doc { ExprDoc value = d->AsDoc<ExprDoc>(eval->value, p->Attr("value")); - if (eval->value->IsInstance<CallNode>()) { + const auto* call = eval->value.as<CallNode>(); + if (call && !call->op.same_as(tirx::builtin::buffer_data())) { return ExprStmtDoc(value); } return ExprStmtDoc(TIR(d, "evaluate")->Call({value}));
diff --git a/src/tirx/transform/flatten_buffer.cc b/src/tirx/transform/flatten_buffer.cc index 72ea51f..8650099 100644 --- a/src/tirx/transform/flatten_buffer.cc +++ b/src/tirx/transform/flatten_buffer.cc
@@ -48,6 +48,9 @@ arith::Analyzer ana; auto pass = BufferFlattener(ana); pass.MarkBufferMapShapes(func); + for (const auto& [param, buffer] : func->buffer_map) { + pass.extern_buffers_.insert(buffer); + } auto body = pass.VisitStmt(func->body); // The buffers in func->buffer_map are deliberately left @@ -121,7 +124,17 @@ } Stmt VisitStmt_(const DeclBufferNode* op) final { - Expr data = VisitExpr(op->data); + Expr data = op->data; + bool is_extern_buffer_source = false; + if (const auto* call = op->data.as<CallNode>(); + call && call->op.same_as(builtin::buffer_data()) && call->args.size() == 1) { + if (const auto* var = call->args[0].as<VarNode>(); var && var->ty.as<BufferTypeNode>()) { + is_extern_buffer_source = extern_buffers_.count(BufferVar(ffi::GetRef<Var>(var))); + } + } + if (!is_extern_buffer_source) { + data = VisitExpr(op->data); + } BufferVar flattened = GetFlattenedBuffer(op->buffer); if (flattened.same_as(op->buffer) && data.same_as(op->data)) { return ffi::GetRef<Stmt>(op); @@ -223,8 +236,8 @@ */ std::unordered_set<BufferVar, ffi::ObjectPtrHash, ffi::ObjectPtrEqual> buffers_used_; - /*! \brief The updated external buffer map. */ - ffi::Map<Var, BufferVar> updated_extern_buffer_map_; + /*! \brief Buffers whose storage is supplied by a PrimFunc parameter. */ + std::unordered_set<BufferVar, ffi::ObjectPtrHash, ffi::ObjectPtrEqual> extern_buffers_; }; PrimFunc FlattenBuffer(PrimFunc f) { return BufferFlattener::Flatten(f); }
diff --git a/src/tirx/transform/lower_intrin.cc b/src/tirx/transform/lower_intrin.cc index db9ca6d..84ac21e 100644 --- a/src/tirx/transform/lower_intrin.cc +++ b/src/tirx/transform/lower_intrin.cc
@@ -42,7 +42,13 @@ namespace tvm { namespace tirx { -static Expr LowerAccessPtr(const CallNode* call) { +struct AccessPtrBufferAlias { + BufferVar buffer; + Expr data; +}; + +static Expr LowerAccessPtr(const CallNode* call, + std::vector<AccessPtrBufferAlias>* buffer_aliases) { TVM_FFI_ICHECK_EQ(call->args.size(), 5U); PrimType dtype = call->args[0].as_or_throw<PrimExpr>().ty(); PrimExpr offset = call->args[2].as_or_throw<PrimExpr>(); @@ -77,21 +83,41 @@ << "tvm_access_ptr expects a buffer Var or nested tvm_access_ptr as args[1], but got " << buffer; Var buffer_var = ffi::GetRef<Var>(buffer_node); + PrimExpr scalar_extent = offset + IntImm(offset.ty(), 1); if (dtype.lanes() != 1) { PrimType offset_ty = offset.ty(); offset = offset * IntImm(offset_ty, dtype.lanes()); + scalar_extent = offset + IntImm(offset_ty, dtype.lanes()); offset = Ramp(offset, IntImm(offset_ty, 1), dtype.lanes()); } + + PrimType scalar_dtype = dtype.WithLanes(1); BufferVar access_buffer{nullptr}; + ffi::String storage_scope; + Expr access_data; if (buffer_var->ty.as<BufferTypeNode>()) { - access_buffer = BufferVar(buffer_var); - TVM_FFI_ICHECK_EQ(access_buffer->dtype, dtype.WithLanes(1)) - << "tvm_access_ptr element type must match the source buffer"; + BufferVar source_buffer(buffer_var); + if (source_buffer->dtype == scalar_dtype && source_buffer->shape.size() == 1) { + access_buffer = source_buffer; + } else { + TVM_FFI_ICHECK_EQ(source_buffer->dtype.WithLanes(1), scalar_dtype) + << "tvm_access_ptr element type must match the source buffer"; + storage_scope = source_buffer.scope(); + access_data = source_buffer.data(); + } } else { auto pointer_type = buffer_var->ty.as_or_throw<PointerType>(); - access_buffer = BufferVar( - buffer_var->name, - BufferType(pointer_type->storage_scope, dtype.WithLanes(1), {offset + 1}, {}, 0, 0, 0)); + storage_scope = pointer_type->storage_scope; + access_data = buffer_var; + } + + if (!access_buffer.defined()) { + // BufferVar identity includes its immutable BufferType. Bind an explicit + // scalar physical view instead of retyping a vector, padded, or packed source. + access_buffer = + BufferVar(buffer_var->name + "_access", + BufferType(storage_scope, scalar_dtype, {scalar_extent}, {}, 0, 0, 0)); + buffer_aliases->push_back({access_buffer, access_data}); } BufferLoad buf_load(access_buffer, {offset}); return Call(call->ty, builtin::address_of(), {buf_load}); @@ -136,9 +162,20 @@ } } + Stmt VisitStmt(const Stmt& stmt) final { + size_t alias_begin = access_ptr_buffer_aliases_.size(); + Stmt result = IRMutatorWithAnalyzer::VisitStmt(stmt); + for (size_t i = access_ptr_buffer_aliases_.size(); i > alias_begin; --i) { + const auto& alias = access_ptr_buffer_aliases_[i - 1]; + result = SeqStmt::Flatten(DeclBuffer(alias.buffer, alias.data), std::move(result)); + } + access_ptr_buffer_aliases_.resize(alias_begin); + return result; + } + Expr VisitExpr_(const CallNode* op) final { if (op->op.same_as(builtin::tvm_access_ptr())) { - return this->VisitExpr(LowerAccessPtr(op)); + return this->VisitExpr(LowerAccessPtr(op, &access_ptr_buffer_aliases_)); } if (auto* ptr_op = op->op.as<OpNode>()) { Op op_ref = ffi::GetRef<Op>(ptr_op); @@ -429,6 +466,7 @@ } std::vector<OpAttrMap<FLowerGeneral>> attr_maps_; + std::vector<AccessPtrBufferAlias> access_ptr_buffer_aliases_; FLowerGeneral fma_{nullptr}; bool support_bitwise_op_{true}; };
diff --git a/src/tirx/transform/vectorize_loop.cc b/src/tirx/transform/vectorize_loop.cc index 81c0318..0cc74d5 100644 --- a/src/tirx/transform/vectorize_loop.cc +++ b/src/tirx/transform/vectorize_loop.cc
@@ -90,6 +90,12 @@ return CheckContains::StmtContains( stmt, [](const PrimExpr& expr) { return expr.as<CallNode>() != nullptr; }); } + +PrimType GetTextureElementType(const Expr& texture) { + const auto* pointer_type = texture->ty.as<PointerTypeNode>(); + TVM_FFI_ICHECK(pointer_type) << "Texture arguments must have PointerType"; + return pointer_type->element_type.as_or_throw<PrimType>(); +} } // namespace inline PrimExpr CreateNewLanes(bool is_scalable, int lanes_or_vscale_factor) { @@ -636,12 +642,8 @@ } else if (op->op.same_as(builtin::texture2d_load())) { int lane = 0; ffi::Array<PrimExpr> fcd = MutateArray({op->args.back().as_or_throw<PrimExpr>()}, &lane); - DLDataType dtype = op->args[0] - .as<VarNode>() - ->ty.as<PointerTypeNode>() - ->element_type.as<PrimTypeNode>() - ->dtype; - TVM_FFI_ICHECK(lane * dtype.bits <= op->args[4].as<IntImmNode>()->value) + PrimType dtype = GetTextureElementType(op->args[0]); + TVM_FFI_ICHECK(lane * dtype.bits() <= op->args[4].as<IntImmNode>()->value) << "Expected Data to be Read is lesser than or equal to Texture Load length"; auto new_args = op->args; @@ -654,12 +656,8 @@ // Vectorize the value to store ffi::Array<PrimExpr> value{op->args.back().as_or_throw<PrimExpr>()}; ffi::Array<PrimExpr> mutated_value = MutateArray(value, &lane); - DLDataType dtype = op->args[0] - .as<VarNode>() - ->ty.as<PointerTypeNode>() - ->element_type.as<PrimTypeNode>() - ->dtype; - TVM_FFI_ICHECK(lane * dtype.bits == op->args[4].as<IntImmNode>()->value) + PrimType dtype = GetTextureElementType(op->args[0]); + TVM_FFI_ICHECK(lane * dtype.bits() == op->args[4].as<IntImmNode>()->value) << "Expected Data to be Written equal to Texture Store length"; ffi::Array<Expr> new_args = op->args; new_args.Set(new_args.size() - 1, mutated_value[0]);
diff --git a/tests/python/all-platform-minimal-test/test_minimal_target_codegen_llvm.py b/tests/python/all-platform-minimal-test/test_minimal_target_codegen_llvm.py index 56ef78c..4b1cdaa 100644 --- a/tests/python/all-platform-minimal-test/test_minimal_target_codegen_llvm.py +++ b/tests/python/all-platform-minimal-test/test_minimal_target_codegen_llvm.py
@@ -54,9 +54,9 @@ dev = tvm.cpu(0) # launch the kernel. n = nn - a = tvm.runtime.tensor(np.random.uniform(size=n).astype(A.dtype), dev) - b = tvm.runtime.tensor(np.random.uniform(size=n).astype(B.dtype), dev) - c = tvm.runtime.tensor(np.zeros(n, dtype=C.dtype), dev) + a = tvm.runtime.tensor(np.random.uniform(size=n).astype(A.dtype.dtype), dev) + b = tvm.runtime.tensor(np.random.uniform(size=n).astype(B.dtype.dtype), dev) + c = tvm.runtime.tensor(np.zeros(n, dtype=C.dtype.dtype), dev) f(a, b, c) tvm.testing.assert_allclose(c.numpy(), a.numpy() + b.numpy())
diff --git a/tests/python/codegen/test_target_codegen.py b/tests/python/codegen/test_target_codegen.py index 7157ae0..b7fa325 100644 --- a/tests/python/codegen/test_target_codegen.py +++ b/tests/python/codegen/test_target_codegen.py
@@ -117,6 +117,32 @@ tvm.compile(func) +@pytest.mark.parametrize( + ("target", "qualifier"), + [("opencl", "__global "), ("metal", "device ")], +) +def test_decl_buffer_offset_preserves_storage_scope(target, qualifier): + @T.prim_func(s_tir=True) + def kernel(A_ptr: T.handle("float32", "global")): + T.func_attr( + { + "calling_conv": 2, + "global_symbol": "kernel", + "tirx.kernel_launch_params": [], + "tirx.noalias": True, + } + ) + A = T.decl_buffer((8,), "float32", data=A_ptr) + B = T.decl_buffer((4,), "float32", data=T.address_of(A[4])) + B[0] = T.float32(1) + + mod = tvm.IRModule({"kernel": kernel}) + build = tvm.get_global_func(f"target.build.{target}") + source = build(mod, tvm.target.Target(target)).inspect_source() + assert f"{qualifier}float* B" in source + assert f"(({qualifier}float*)" in source + + @pytest.mark.parametrize("target", ["c", "llvm"]) def test_codegen_loop_step(target): if target != "c" and not tvm.testing.device_enabled(target):
diff --git a/tests/python/codegen/test_target_codegen_c_host.py b/tests/python/codegen/test_target_codegen_c_host.py index f33eb92..989dc21 100644 --- a/tests/python/codegen/test_target_codegen_c_host.py +++ b/tests/python/codegen/test_target_codegen_c_host.py
@@ -245,6 +245,33 @@ built.export_library(temp.relpath("workspace.so")) +def test_local_alloc_buffer_uses_plain_c_pointer(): + @I.ir_module(s_tir=True) + class Module: + @T.prim_func(s_tir=True) + def main(A: T.Buffer((1,), "float32")): + B = T.alloc_buffer((1,), "float32", scope="local") + for i in range(1): + with T.sblock("copy"): + vi = T.axis.spatial(1, i) + T.reads(A[vi]) + T.writes(B[vi], A[vi]) + B[vi] = A[vi] + T.float32(1) + A[vi] = B[vi] + + built = tvm.tirx.build(Module, target="c") + assert "local float*" not in built.inspect_source() + + temp = utils.tempdir() + path_dso = temp.relpath("local_alloc.so") + built.export_library(path_dso) + loaded = tvm.runtime.load_module(path_dso) + + data = tvm.runtime.tensor(np.array([1.0], dtype="float32")) + loaded["main"](data) + tvm.testing.assert_allclose(data.numpy(), np.array([2.0], dtype="float32")) + + def test_vector_access_ptr_address_uses_ramp_base(): buffer = tvm.tirx.decl_buffer((8,), "float32x2", name="A") access_ptr = buffer.access_ptr(access_mask=3, offset=2, extent=4)
diff --git a/tests/python/codegen/test_target_codegen_cuda.py b/tests/python/codegen/test_target_codegen_cuda.py index 7fbcb89..c921393 100644 --- a/tests/python/codegen/test_target_codegen_cuda.py +++ b/tests/python/codegen/test_target_codegen_cuda.py
@@ -1015,16 +1015,16 @@ for blockIdx in T.thread_binding(1, thread="blockIdx.x"): for threadIdx in T.thread_binding(128, thread="threadIdx.x"): if threadIdx == 0: - A[0, 0] = T.reinterpret("float64", A_map) + A[0, 0] = T.Cast("float32", T.address_of(A_map)) # fmt: on mod = tvm.IRModule({"main": main}) mod = tvm.compile(mod, target="cuda") assert ( """ -extern "C" __global__ void __launch_bounds__(128) main_kernel(const __grid_constant__ CUtensorMap A_map, float* __restrict__ A_ptr) { +extern "C" __global__ void __launch_bounds__(128) main_kernel(float* __restrict__ A_ptr, const __grid_constant__ CUtensorMap A_map) { if (((int)threadIdx.x) == 0) { - A_ptr[0] = ((float)(*(double *)(&(A_map)))); + A_ptr[0] = ((float)((unsigned long long)(&(A_map)))); } }""".strip() in mod.mod.imports[0].inspect_source()
diff --git a/tests/python/codegen/test_target_codegen_llvm.py b/tests/python/codegen/test_target_codegen_llvm.py index ed53646..6f9f128 100644 --- a/tests/python/codegen/test_target_codegen_llvm.py +++ b/tests/python/codegen/test_target_codegen_llvm.py
@@ -1263,6 +1263,22 @@ tvm.compile(Module) +def test_invalid_volatile_masked_decl_buffer_load(): + @I.ir_module(s_tir=True) + class Module: + @T.prim_func(s_tir=True) + def main(b: T.handle): + B = T.match_buffer(b, [4]) + A = T.alloc_buffer((4,), annotations={"tirx.volatile": True}) + A_alias = T.decl_buffer((4,), data=A.data) + B[0:4] = A_alias.vload([T.Ramp(0, 1, 4)], predicate=T.Broadcast(T.bool(True), 4)) + + err_msg = "The masked load intrinsic does not support declaring load as volatile." + with pytest.raises(RuntimeError, match=err_msg): + with tvm.target.Target("llvm"): + tvm.compile(Module) + + def test_invalid_volatile_masked_buffer_store(): @I.ir_module(s_tir=True) class Module:
diff --git a/tests/python/codegen/test_target_codegen_metal.py b/tests/python/codegen/test_target_codegen_metal.py index 223c670..a1fb713 100644 --- a/tests/python/codegen/test_target_codegen_metal.py +++ b/tests/python/codegen/test_target_codegen_metal.py
@@ -349,5 +349,40 @@ host_lib.export_library(lib_path) +def test_codegen_simdgroup_buffer_data(): + """Simdgroup intrinsics should accept buffer_data projections.""" + + @I.ir_module(s_tir=True) + class Module: + @T.prim_func(s_tir=True) + def kernel(): + T.func_attr( + { + "calling_conv": 2, + "global_symbol": "kernel", + "tirx.kernel_launch_params": [], + } + ) + A = T.alloc_buffer((64,), "float16", scope="shared") + A_frag = T.alloc_buffer((64,), "float16", scope="metal.simdgroup") + B_frag = T.alloc_buffer((64,), "float16", scope="metal.simdgroup") + C_frag = T.alloc_buffer((64,), "float16", scope="metal.simdgroup") + T.metal.make_filled_simdgroup_matrix(C_frag.data, 0, T.float32(0), 8, 8) + T.metal.simdgroup_load(A_frag.data, 0, A.data, 8, 8, 8, T.bool(False)) + T.metal.simdgroup_store(C_frag.data, 0, A.data, 8, 8, 8, T.bool(False)) + T.metal.simdgroup_multiply_accumulate( + C_frag.data, 0, A_frag.data, 0, B_frag.data, 0, C_frag.data, 0 + ) + + metal_codegen = tvm.get_global_func("target.build.metal") + module = metal_codegen(Module, tvm.target.Target("metal")) + source = module.inspect_source() + + assert "make_filled_simdgroup_matrix<half, 8, 8>" in source + assert "simdgroup_load(" in source + assert "simdgroup_store(" in source + assert "simdgroup_multiply_accumulate(" in source + + if __name__ == "__main__": tvm.testing.main()
diff --git a/tests/python/codegen/test_target_codegen_vulkan.py b/tests/python/codegen/test_target_codegen_vulkan.py index 7ff86da..d3213b4 100644 --- a/tests/python/codegen/test_target_codegen_vulkan.py +++ b/tests/python/codegen/test_target_codegen_vulkan.py
@@ -572,19 +572,32 @@ @pytest.mark.gpu @pytest.mark.skipif(not env.has_vulkan(), reason="need vulkan") def test_codegen_decl_buffer(): - """The codegen should accept DeclBuffer nodes in its input""" + """DeclBuffer aliases should retain their backing storage metadata.""" @I.ir_module(s_tir=True) - class Module: + class AllocationBacked: @T.prim_func(s_tir=True) def kernel(): T.func_attr({"calling_conv": 2, "global_symbol": "kernel", "tirx.noalias": True}) A = T.alloc_buffer((256,), dtype="float32", scope="local") A_buf = T.decl_buffer([256], dtype="float32", scope="local", data=A.data) + A_buf[0] = T.float32(1) + T.evaluate(A_buf[0]) target = tvm.target.Target("vulkan") vulkan_codegen = tvm.get_global_func("target.build.vulkan") - vulkan_codegen(Module, target) + vulkan_codegen(AllocationBacked, target) + + @I.ir_module(s_tir=True) + class ParameterBacked: + @T.prim_func(s_tir=True) + def main(A: T.Buffer((1,), "float32"), B: T.Buffer((1,), "float32")): + A_buf = T.decl_buffer([1], dtype="float32", data=A.data) + B_buf = T.decl_buffer([1], dtype="float32", data=B.data) + for tx in T.thread_binding(1, thread="threadIdx.x"): + B_buf[tx] = A_buf[tx] + + tvm.compile(ParameterBacked, target=target) @pytest.mark.gpu
diff --git a/tests/python/relax/test_blockbuilder_core.py b/tests/python/relax/test_blockbuilder_core.py index bdee939..8df1ca8 100644 --- a/tests/python/relax/test_blockbuilder_core.py +++ b/tests/python/relax/test_blockbuilder_core.py
@@ -389,7 +389,7 @@ buffer_B = f_matmul.buffer_map[param_B] assert param_A.name != param_B.name assert buffer_A.name != buffer_B.name - assert buffer_A.data.name != buffer_B.data.name + assert not buffer_A.same_as(buffer_B) def test_call_te_with_unsupported_shape_arg():
diff --git a/tests/python/runtime/test_runtime_trace.py b/tests/python/runtime/test_runtime_trace.py index 2ab8dbe..c2ddb04 100644 --- a/tests/python/runtime/test_runtime_trace.py +++ b/tests/python/runtime/test_runtime_trace.py
@@ -26,8 +26,8 @@ x = te.placeholder((n, n, n), name="X", dtype="float32") y = te.compute(x.shape, lambda i, j, k: tvm.tirx.trace([i, j, k, x[i][j][k]])) f = tvm.compile(te.create_prim_func([x, y]), target="llvm") - xnd = tvm.runtime.tensor(np.ones((n, n, n), dtype=x.dtype)) - ynd = tvm.runtime.tensor(np.zeros((n, n, n), dtype=y.dtype)) + xnd = tvm.runtime.tensor(np.ones((n, n, n), dtype=x.dtype.dtype)) + ynd = tvm.runtime.tensor(np.zeros((n, n, n), dtype=y.dtype.dtype)) f(xnd, ynd) @@ -47,9 +47,9 @@ ) f = tvm.compile(te.create_prim_func([x, y, z]), "llvm") - xnd = tvm.runtime.tensor(np.ones((n, n, n), dtype=x.dtype)) - ynd = tvm.runtime.tensor(np.zeros((n, n, n), dtype=y.dtype)) - znd = tvm.runtime.tensor(np.zeros((n, n, n), dtype=z.dtype)) + xnd = tvm.runtime.tensor(np.ones((n, n, n), dtype=x.dtype.dtype)) + ynd = tvm.runtime.tensor(np.zeros((n, n, n), dtype=y.dtype.dtype)) + znd = tvm.runtime.tensor(np.zeros((n, n, n), dtype=z.dtype.dtype)) f(xnd, ynd, znd) assert np.array_equal(xnd.numpy(), np.ones((n, n, n))) @@ -77,9 +77,9 @@ ), ) f = tvm.compile(te.create_prim_func([a, b, c])) - xnd = tvm.runtime.tensor(np.array(np.ones((n, n, n), dtype=a.dtype))) - ynd = tvm.runtime.tensor(np.array(np.ones((n, n, n), dtype=b.dtype))) - znd = tvm.runtime.tensor(np.zeros((n, n, n), dtype=c.dtype)) + xnd = tvm.runtime.tensor(np.array(np.ones((n, n, n), dtype=a.dtype.dtype))) + ynd = tvm.runtime.tensor(np.array(np.ones((n, n, n), dtype=b.dtype.dtype))) + znd = tvm.runtime.tensor(np.zeros((n, n, n), dtype=c.dtype.dtype)) f(xnd, ynd, znd) assert np.array_equal(znd.numpy(), xnd.numpy() + ynd.numpy()) @@ -109,11 +109,11 @@ ), ) f = tvm.compile(te.create_prim_func([a, b, d, e, c])) - a_nd = tvm.runtime.tensor(np.array(np.ones((n, n, n), dtype=a.dtype))) - b_nd = tvm.runtime.tensor(np.array(np.ones((n, n, n), dtype=b.dtype))) - d_nd = tvm.runtime.tensor(np.array(np.ones((n, n, n), dtype=d.dtype))) - e_nd = tvm.runtime.tensor(np.array(np.ones((n, n, n), dtype=e.dtype))) - c_nd = tvm.runtime.tensor(np.zeros((n, n, n), dtype=c.dtype)) + a_nd = tvm.runtime.tensor(np.array(np.ones((n, n, n), dtype=a.dtype.dtype))) + b_nd = tvm.runtime.tensor(np.array(np.ones((n, n, n), dtype=b.dtype.dtype))) + d_nd = tvm.runtime.tensor(np.array(np.ones((n, n, n), dtype=d.dtype.dtype))) + e_nd = tvm.runtime.tensor(np.array(np.ones((n, n, n), dtype=e.dtype.dtype))) + c_nd = tvm.runtime.tensor(np.zeros((n, n, n), dtype=c.dtype.dtype)) f(a_nd, b_nd, d_nd, e_nd, c_nd) assert np.array_equal( c_nd.numpy(), a_nd.numpy() + b_nd.numpy() + d_nd.numpy() + e_nd.numpy() @@ -140,11 +140,17 @@ ), ) f = tvm.compile(te.create_prim_func([a, b, c])) - npa = np.array([[1, 0, 0, 0], [0, 1, 0, 0], [0, 0, 1, 0], [0, 0, 0, 1]], dtype=a.dtype) - npb = np.array([[1, 0, 0, 0], [0, 1, 0, 0], [0, 0, 1, 0], [0, 0, 0, 1]], dtype=a.dtype) + npa = np.array( + [[1, 0, 0, 0], [0, 1, 0, 0], [0, 0, 1, 0], [0, 0, 0, 1]], + dtype=a.dtype.dtype, + ) + npb = np.array( + [[1, 0, 0, 0], [0, 1, 0, 0], [0, 0, 1, 0], [0, 0, 0, 1]], + dtype=a.dtype.dtype, + ) xnd = tvm.runtime.tensor(npa) ynd = tvm.runtime.tensor(npb) - znd = tvm.runtime.tensor(np.zeros((n, n), dtype=c.dtype)) + znd = tvm.runtime.tensor(np.zeros((n, n), dtype=c.dtype.dtype)) f(xnd, ynd, znd) assert np.array_equal(znd.numpy(), npa + npb) @@ -170,9 +176,9 @@ ) f = tvm.compile(te.create_prim_func([x, y, z])) - xnd = tvm.runtime.tensor(np.ones((n,), dtype=x.dtype)) - ynd = tvm.runtime.tensor(np.zeros((n,), dtype=y.dtype)) - znd = tvm.runtime.tensor(np.zeros((n,), dtype=z.dtype)) + xnd = tvm.runtime.tensor(np.ones((n,), dtype=x.dtype.dtype)) + ynd = tvm.runtime.tensor(np.zeros((n,), dtype=y.dtype.dtype)) + znd = tvm.runtime.tensor(np.zeros((n,), dtype=z.dtype.dtype)) f(xnd, ynd, znd) check_array_first = np.array([13, 13, 13, 13]) check_array_second = np.array([14, 14, 14, 14]) @@ -203,9 +209,9 @@ ) f = tvm.compile(te.create_prim_func([x, y, z]), target="llvm") - xnd = tvm.runtime.tensor(np.ones((n,), dtype=x.dtype)) - ynd = tvm.runtime.tensor(np.zeros((n,), dtype=y.dtype)) - znd = tvm.runtime.tensor(np.zeros((n,), dtype=z.dtype)) + xnd = tvm.runtime.tensor(np.ones((n,), dtype=x.dtype.dtype)) + ynd = tvm.runtime.tensor(np.zeros((n,), dtype=y.dtype.dtype)) + znd = tvm.runtime.tensor(np.zeros((n,), dtype=z.dtype.dtype)) f(xnd, ynd, znd) check_array_first = np.array([13.0, 13.0, 13.0, 13.0]) check_array_second = np.array([14.0, 14.0, 14.0, 14.0])
diff --git a/tests/python/s_tir/analysis/test_sblock_access_region.py b/tests/python/s_tir/analysis/test_sblock_access_region.py index 16c4202..daf9430 100644 --- a/tests/python/s_tir/analysis/test_sblock_access_region.py +++ b/tests/python/s_tir/analysis/test_sblock_access_region.py
@@ -121,6 +121,18 @@ @T.prim_func(s_tir=True) +def decl_buffer_alias_func( + A: T.Buffer((16,), "float32"), + B: T.Buffer((16,), "float32"), +) -> None: + with T.sblock("alias"): + T.reads(A[0]) + T.writes(B[0]) + A_view = T.decl_buffer((16,), "float32", data=A.data) + B[0] = A[0] + A_view[0] + + +@T.prim_func(s_tir=True) def access_in_if_then_else_func() -> None: A = T.sblock_alloc_buffer([8]) B = T.sblock_alloc_buffer([8]) @@ -263,6 +275,16 @@ tvm.ir.assert_structural_equal(ret0[1], ret1[1]) +def test_decl_buffer_alias_is_not_an_opaque_access(): + block = decl_buffer_alias_func.body.block + buffer_var_map = {buf: buf for buf in decl_buffer_alias_func.buffer_map.values()} + + reads, writes, opaque = s_tir.analysis.get_sblock_access_region(block, buffer_var_map) + tvm.ir.assert_structural_equal(block.reads, reads) + tvm.ir.assert_structural_equal(block.writes, writes) + tvm.ir.assert_structural_equal([], opaque) + + def test_match_buffer(): root_block = match_buffer_func.body.block block = root_block.body.body.body.block
diff --git a/tests/python/s_tir/transform/test_s_tir_transform_lower_cross_thread_reduction.py b/tests/python/s_tir/transform/test_s_tir_transform_lower_cross_thread_reduction.py index 2284bf3..88e4c7e 100644 --- a/tests/python/s_tir/transform/test_s_tir_transform_lower_cross_thread_reduction.py +++ b/tests/python/s_tir/transform/test_s_tir_transform_lower_cross_thread_reduction.py
@@ -1901,6 +1901,15 @@ _check(thread_broadcast_2, lowered_thread_broadcast_2) +def test_thread_broadcast_rewrite_2_full_pipeline(): + target_host = "llvm" if tvm.runtime.enabled("llvm") else "c" + target = tvm.target.Target("cuda").with_host(target_host) + mod = tvm.IRModule.from_expr(thread_broadcast_2.with_attr("global_symbol", "main")) + mod = tvm.tirx.transform.BindTarget(target)(mod) + pipeline, _, _ = tvm.tirx.get_tir_pipeline("s_tir") + pipeline(mod) + + def test_no_thread_broadcast_rewrite(): _check(no_thread_broadcast, lowered_no_thread_broadcast)
diff --git a/tests/python/s_tir/transform/test_s_tir_transform_lower_thread_all_reduce.py b/tests/python/s_tir/transform/test_s_tir_transform_lower_thread_all_reduce.py index e334efb..5c08e66 100644 --- a/tests/python/s_tir/transform/test_s_tir_transform_lower_thread_all_reduce.py +++ b/tests/python/s_tir/transform/test_s_tir_transform_lower_thread_all_reduce.py
@@ -194,6 +194,86 @@ assert "tvm_warp_shuffle" in After_script +def test_multi_group_reduction_consumed_through_alias(): + transform = tvm.s_tir.transform.LowerThreadAllreduce() + + @I.ir_module + class Before: + @T.prim_func(private=True, s_tir=True) + def main(A: T.Buffer((4, 128), "float32"), B: T.Buffer((4,), "float32")): + T.func_attr({"target": T.target("cuda", host="llvm")}) + threadIdx_y = T.launch_thread("threadIdx.y", 4) + cross_thread_B = T.alloc_buffer((1,), scope="local") + threadIdx_x = T.launch_thread("threadIdx.x", 128) + cross_thread_B_alias = T.decl_buffer((1,), data=cross_thread_B.data, scope="local") + with T.attr( + T.comm_reducer(lambda x0, y0: x0 + y0, [T.float32(0)]), + "reduce_scope", + T.int32(0), + ): + A_flat = T.decl_buffer((512,), data=A.data) + T.tvm_thread_allreduce( + T.uint32(1), + A_flat[threadIdx_y * 128 + threadIdx_x], + T.bool(True), + cross_thread_B[0], + threadIdx_x, + ) + cross_thread_B_alias[0] = cross_thread_B[0] + if threadIdx_x == 0: + B_flat = T.decl_buffer((4,), data=B.data) + B_flat[threadIdx_y] = cross_thread_B_alias[0] + + After = transform(Before) + assert tvm.tirx.analysis.verify_well_formed(After) + After_script = After.script() + assert "red_result[threadIdx_y] = red_result[threadIdx_y]" in After_script + assert "B_flat[threadIdx_y] = red_result[threadIdx_y]" in After_script + assert "B_flat[threadIdx_y] = red_result[0]" not in After_script + assert "red_result[0] = red_result[threadIdx_y]" not in After_script + assert "cross_thread_B_alias" not in After_script + + +def test_multi_group_reduction_with_alias_declared_after_allreduce(): + transform = tvm.s_tir.transform.LowerThreadAllreduce() + + @I.ir_module + class Before: + @T.prim_func(private=True, s_tir=True) + def main(A: T.Buffer((4, 128), "float32"), B: T.Buffer((4,), "float32")): + T.func_attr({"target": T.target("cuda", host="llvm")}) + threadIdx_y = T.launch_thread("threadIdx.y", 4) + cross_thread_B = T.alloc_buffer((1,), scope="local") + threadIdx_x = T.launch_thread("threadIdx.x", 128) + with T.attr( + T.comm_reducer(lambda x0, y0: x0 + y0, [T.float32(0)]), + "reduce_scope", + T.int32(0), + ): + A_flat = T.decl_buffer((512,), data=A.data) + T.tvm_thread_allreduce( + T.uint32(1), + A_flat[threadIdx_y * 128 + threadIdx_x], + T.bool(True), + cross_thread_B[0], + threadIdx_x, + ) + cross_thread_B_alias = T.decl_buffer((1,), data=cross_thread_B.data, scope="local") + cross_thread_B[0] = cross_thread_B_alias[0] + if threadIdx_x == 0: + B_flat = T.decl_buffer((4,), data=B.data) + B_flat[threadIdx_y] = cross_thread_B_alias[0] + + After = transform(Before) + assert tvm.tirx.analysis.verify_well_formed(After) + After_script = After.script() + assert "red_result[threadIdx_y] = red_result[threadIdx_y]" in After_script + assert "B_flat[threadIdx_y] = red_result[threadIdx_y]" in After_script + assert "B_flat[threadIdx_y] = red_result[0]" not in After_script + assert "red_result[0] = red_result[threadIdx_y]" not in After_script + assert "cross_thread_B_alias" not in After_script + + def test_multi_group_mask1(): transform = tvm.s_tir.transform.LowerThreadAllreduce()
diff --git a/tests/python/s_tir/transform/test_s_tir_transform_thread_sync.py b/tests/python/s_tir/transform/test_s_tir_transform_thread_sync.py index 9aeb7df..8488f2f 100644 --- a/tests/python/s_tir/transform/test_s_tir_transform_thread_sync.py +++ b/tests/python/s_tir/transform/test_s_tir_transform_thread_sync.py
@@ -102,6 +102,29 @@ tvm.ir.assert_structural_equal(mod["main"], expected) +def test_sync_shared_aliasing_buffer_views(): + @T.prim_func(private=True, s_tir=True) + def func(A: T.Buffer((64,), "float32")): + blockIdx_x = T.launch_thread("blockIdx.x", 1) + shared_storage = T.alloc_buffer((32,), "float16", scope="shared") + local = T.alloc_buffer((1,), "float32", scope="local") + threadIdx_x = T.launch_thread("threadIdx.x", 32) + shared_half = T.decl_buffer((32,), "float16", data=shared_storage.data, scope="shared") + shared_float = T.decl_buffer((16,), "float32", data=shared_storage.data, scope="shared") + for i in range(2): + shared_half[threadIdx_x] = T.Cast("float16", A[i * 32 + threadIdx_x]) + T.tvm_storage_sync("shared") + local[0] = shared_float[threadIdx_x % 16] + A[i * 32 + threadIdx_x] = local[0] + + mod = tvm.IRModule({"main": func}) + mod = tvm.s_tir.transform.ThreadSync("shared")(mod) + + # In addition to the explicit write-to-read barrier, the shared physical + # storage needs a read-to-next-write barrier across loop iterations. + assert str(mod["main"]).count("T.tvm_storage_sync") == 2 + + @pytest.mark.gpu @pytest.mark.skipif(not env.has_cuda(), reason="need cuda") def test_sync_bind():
diff --git a/tests/python/tirx-base/test_tir_texture_scope.py b/tests/python/tirx-base/test_tir_texture_scope.py index dc90008..312f3f7 100644 --- a/tests/python/tirx-base/test_tir_texture_scope.py +++ b/tests/python/tirx-base/test_tir_texture_scope.py
@@ -14,7 +14,7 @@ # KIND, either express or implied. See the License for the # specific language governing permissions and limitations # under the License. -# ruff: noqa: F401, F841 +# ruff: noqa: F401 import pytest @@ -22,7 +22,10 @@ import tvm.testing from tvm import tirx from tvm.ir.module import IRModule +from tvm.s_tir.backend.adreno import pipeline as adreno_pipeline from tvm.script import tirx as T +from tvm.tirx.build import split_host_device_mods +from tvm.tirx.compilation_pipeline import finalize_device_passes def test_texture_scope(): @@ -56,8 +59,18 @@ schedule_block(sch.get_sblock("B")) schedule_block(sch.get_sblock("C")) - target = tvm.target.Target("opencl") - mod = tvm.compile(sch.mod["main"], target=target) + target = tvm.target.Target({"kind": "opencl", "keys": ["adreno"]}) + lowered = tirx.transform.BindTarget(target.with_host("c"))(sch.mod) + lowered = adreno_pipeline.default_tir_pipeline()(lowered) + _, device_mods = split_host_device_mods(lowered) + assert len(device_mods) == 1 + device_target, device_mod = next(iter(device_mods.items())) + device_mod = finalize_device_passes()(device_mod) + source = tvm.get_global_func("target.build.opencl")(device_mod, device_target).inspect_source() + assert "__read_only image2d_array_t" in source + assert "__write_only image2d_array_t" in source + assert "READ_IMAGEF" in source + assert "write_imagef" in source if __name__ == "__main__":
diff --git a/tests/python/tirx-transform/test_tir_transform_convert_ssa.py b/tests/python/tirx-transform/test_tir_transform_convert_ssa.py index 59a9a5c..eb7166a 100644 --- a/tests/python/tirx-transform/test_tir_transform_convert_ssa.py +++ b/tests/python/tirx-transform/test_tir_transform_convert_ssa.py
@@ -545,14 +545,19 @@ elem_offset=loop_var * 128, scope="shared.dyn", ) + buffer_data = tirx.Var("buffer_data", buffer.data.ty) loop = tirx.For( loop_var, 0, 128, tirx.ForKind.SERIAL, - tirx.DeclBuffer(buffer, tirx.Evaluate(tirx.BufferLoad(buffer, [0]))), + tirx.DeclBuffer( + buffer, + tirx.Evaluate(tirx.BufferLoad(buffer, [0])), + data=buffer_data, + ), ) - func = tirx.PrimFunc([buffer.data], tirx.SeqStmt([loop, loop, loop])) + func = tirx.PrimFunc([buffer_data], tirx.SeqStmt([loop, loop, loop])) after = tvm.tirx.transform.ConvertSSA()(tvm.IRModule.from_expr(func))
diff --git a/tests/python/tirx-transform/test_tir_transform_lower_intrin.py b/tests/python/tirx-transform/test_tir_transform_lower_intrin.py index dda7066..ccfa048 100644 --- a/tests/python/tirx-transform/test_tir_transform_lower_intrin.py +++ b/tests/python/tirx-transform/test_tir_transform_lower_intrin.py
@@ -113,6 +113,19 @@ assert isinstance(load, tvm.tirx.BufferLoad) assert int(tvm.arith.Analyzer().simplify(load.indices[0])) == 5 + targets = ["c"] + if env.has_llvm(): + targets.append("llvm") + for target in targets: + target = tvm.target.Target(target) + build_func = ( + tvm.tirx.PrimFunc([data], body) + .with_attr("global_symbol", "main") + .with_attr("target", target) + ) + build_mod = tvm.tirx.transform.LowerIntrin()(tvm.IRModule.from_expr(build_func)) + tvm.tirx.build(build_mod, target=target) + def test_lower_vector_access_ptr(): buffer = tvm.tirx.decl_buffer((8,), "float32x2", name="A") @@ -128,13 +141,20 @@ "target", tvm.target.Target("llvm") ) ) - lowered = tvm.tirx.transform.LowerIntrin()(mod)["main"].body.value + lowered_body = tvm.tirx.transform.LowerIntrin()(mod)["main"].body + assert isinstance(lowered_body, tvm.tirx.SeqStmt) + alias = lowered_body.seq[0] + assert isinstance(alias, tvm.tirx.DeclBuffer) + assert alias.data.op.name == "tirx.buffer_data" + assert alias.data.args[0].same_as(buffer) + lowered = lowered_body.seq[1].value assert lowered.op.name == "tirx.address_of" assert lowered.ty == access_ptr.ty load = lowered.args[0] assert isinstance(load, tvm.tirx.BufferLoad) - assert load.buffer.same_as(buffer) + assert load.buffer.same_as(alias.buffer) + assert not load.buffer.same_as(buffer) assert load.buffer.ty.dtype == tvm.ir.PrimType("float32") assert len(load.indices) == 1 ramp = load.indices[0] @@ -170,6 +190,25 @@ assert int(load.indices[0]) == 3 +@pytest.mark.parametrize("shape", [(), (2, 4)]) +def test_lower_access_ptr_uses_flat_alias_for_non_1d_buffer(shape): + buffer = tvm.tirx.decl_buffer(shape, "float32", "buffer") + access = buffer.access_ptr(access_mask=1) + func = tvm.tirx.PrimFunc([buffer], tvm.tirx.Evaluate(access)).with_attr( + "target", tvm.target.Target("llvm") + ) + + lowered = tvm.tirx.transform.LowerIntrin()(tvm.IRModule.from_expr(func))["main"].body + assert isinstance(lowered, tvm.tirx.SeqStmt) + alias = lowered.seq[0] + assert isinstance(alias, tvm.tirx.DeclBuffer) + assert len(alias.buffer.ty.shape) == 1 + load = lowered.seq[1].value.args[0] + assert isinstance(load, tvm.tirx.BufferLoad) + assert load.buffer.same_as(alias.buffer) + assert len(load.indices) == 1 + + def get_ref_data(): """Get reference data for every pairs""" import itertools
diff --git a/tests/python/tirx/codegen/test_codegen_cuda.py b/tests/python/tirx/codegen/test_codegen_cuda.py index 745d10d..1b31c2b 100644 --- a/tests/python/tirx/codegen/test_codegen_cuda.py +++ b/tests/python/tirx/codegen/test_codegen_cuda.py
@@ -45,11 +45,17 @@ def test_vector_access_ptr_preserves_packed_offset(monkeypatch): buffer = tvm.tirx.decl_buffer((8,), "int4x4", name="A") + data = tvm.tirx.Var("A_data", tvm.tirx.buffer_data_pointer_type(buffer)) access_ptr = buffer.access_ptr(access_mask=3, offset=2, extent=4) - body = tvm.tirx.Evaluate(tvm.tirx.call_extern("void", "consume", access_ptr)) + body = tvm.tirx.SeqStmt( + [ + tvm.tirx.DeclBuffer(buffer, data=data), + tvm.tirx.Evaluate(tvm.tirx.call_extern("void", "consume", access_ptr)), + ] + ) target = tvm.target.Target({"kind": "cuda", "arch": "sm_80"}) func = ( - tvm.tirx.PrimFunc([buffer.data], body) + tvm.tirx.PrimFunc([data], body) .with_attr("global_symbol", "main") .with_attr("target", target) )
diff --git a/tests/python/tirx/codegen/test_codegen_dsmem.py b/tests/python/tirx/codegen/test_codegen_dsmem.py index 1a960d7..1905d55 100644 --- a/tests/python/tirx/codegen/test_codegen_dsmem.py +++ b/tests/python/tirx/codegen/test_codegen_dsmem.py
@@ -106,11 +106,14 @@ # fmt: on binds = [] + decl_buffers = [] loads = [] def collect(node): if isinstance(node, tvm.tirx.Bind): binds.append(node) + elif isinstance(node, tvm.tirx.DeclBuffer): + decl_buffers.append(node) elif isinstance(node, tvm.tirx.BufferLoad): loads.append(node) @@ -120,11 +123,14 @@ assert binds[0].var.ty.storage_scope == "shared" assert binds[0].value.ty.storage_scope == "shared" assert_structural_equal(binds[0].var.ty, binds[0].value.ty) - assert any(load.buffer.data.same_as(binds[0].var) for load in loads) + assert len(decl_buffers) == 1 + assert decl_buffers[0].data.same_as(binds[0].var) + assert any(load.buffer.same_as(decl_buffers[0].buffer) for load in loads) assert_structural_equal(main, tvm.script.from_source(main.script())) src = _get_source(main) - assert "uint64_t* remote_mbar_ptr" in src + assert "uint64_t* remote_ptr" in src + assert "A_ptr[0] = remote_ptr[0]" in src assert "tvm_builtin_ptx_mapa_u64" in src
diff --git a/tests/python/tirx/test_op.py b/tests/python/tirx/test_op.py index dde3662..efd245d 100644 --- a/tests/python/tirx/test_op.py +++ b/tests/python/tirx/test_op.py
@@ -62,10 +62,12 @@ ) restored = pickle.loads(pickle.dumps(call)) - assert_structural_equal(restored, call) + # Buffer values are ordinary free Vars. Pickle reconstructs their + # identities, so compare them under the standard free-Var mapping. + assert_structural_equal(restored, call, map_free_vars=True) assert restored.op.same_as(call.op) - assert_structural_equal(restored.args, call.args) - assert_structural_equal(restored.workspace, call.workspace) + assert_structural_equal(restored.args, call.args, map_free_vars=True) + assert_structural_equal(restored.workspace, call.workspace, map_free_vars=True) assert_structural_equal(restored.config, call.config) assert restored.dispatch == call.dispatch assert_structural_equal(restored.scope, call.scope)
diff --git a/tests/python/tvmscript/test_tvmscript_printer_structural_equal.py b/tests/python/tvmscript/test_tvmscript_printer_structural_equal.py index 6441cd6..f25f53c 100644 --- a/tests/python/tvmscript/test_tvmscript_printer_structural_equal.py +++ b/tests/python/tvmscript/test_tvmscript_printer_structural_equal.py
@@ -76,12 +76,14 @@ AccessPath.root() .attr("buffer_map") .map_item(func1.params[1]) + .attr("ty") .attr("shape") .array_item(1) .attr("value"), AccessPath.root() .attr("buffer_map") .map_item(func2.params[1]) + .attr("ty") .attr("shape") .array_item(1) .attr("value"), @@ -139,8 +141,20 @@ assert _error_message(ve.value) == _expected_result( func1, func2, - AccessPath.root().attr("body").attr("buffer").attr("shape").array_item(0).attr("value"), - AccessPath.root().attr("body").attr("buffer").attr("shape").array_item(0).attr("value"), + AccessPath.root() + .attr("body") + .attr("buffer") + .attr("ty") + .attr("shape") + .array_item(0) + .attr("value"), + AccessPath.root() + .attr("body") + .attr("buffer") + .attr("ty") + .attr("shape") + .array_item(0) + .attr("value"), )
diff --git a/tests/python/tvmscript/test_tvmscript_printer_tir.py b/tests/python/tvmscript/test_tvmscript_printer_tir.py index c29b99e..ac0d761 100644 --- a/tests/python/tvmscript/test_tvmscript_printer_tir.py +++ b/tests/python/tvmscript/test_tvmscript_printer_tir.py
@@ -90,7 +90,7 @@ ) -def test_prim_func_no_sugar_shared_buffer_data(): +def test_prim_func_buffer_data_argument_is_scope_hint(): a = tirx.Var("a", "handle") b = tirx.Var("b", "handle") buffer_data = tirx.decl_buffer(shape=[128, 128], dtype="float32", name="A").data @@ -114,9 +114,7 @@ # from tvm.tirx.layout import Axis @T.prim_func(s_tir=True) -def main(a: T.handle, b: T.handle): - A = T.match_buffer(a, (128, 128)) - B = T.match_buffer(b, (256, 256), data=A.data) +def main(A: T.Buffer((128, 128), "float32"), B: T.Buffer((256, 256), "float32")): T.evaluate(0) """, )