diff --git a/src/target/llvm/codegen_aarch64.cc b/src/target/llvm/codegen_aarch64.cc index 414b9a75c4d3..3c62ca095cfc 100644 --- a/src/target/llvm/codegen_aarch64.cc +++ b/src/target/llvm/codegen_aarch64.cc @@ -40,18 +40,68 @@ class CodeGenAArch64 final : public CodeGenCPU { CodeGenAArch64() = default; virtual ~CodeGenAArch64() = default; - void Dispatch_(const AttrStmtNode* op); - void AddFunction(const GlobalVar& gvar, const PrimFunc& f); - void SetTargetAttributes(llvm::Function* func); + llvm::Function* DeclareFunction(const GlobalVar& gvar, const PrimFunc& f) override; + void AddFunction(const GlobalVar& gvar, const PrimFunc& f) override; + void SetTargetAttributes(llvm::Function* func) override; - bool func_has_pstate_sm = false; - bool func_has_pstate_za = false; + private: + void SetComputeScopeAttributes(llvm::Function* func) override; + void SetPStateAttributes(llvm::Function* func, const PrimFunc& f); + + ffi::Optional current_prim_func_; + llvm::Function* packed_function_{nullptr}; }; +llvm::Function* CodeGenAArch64::DeclareFunction(const GlobalVar& gvar, const PrimFunc& f) { + llvm::Function* func = CodeGenCPU::DeclareFunction(gvar, f); + if (f->GetAttr(tvm::attr::kCallingConv) != CallingConv::kCPackedFunc) { + SetPStateAttributes(func, f); + } + return func; +} + void CodeGenAArch64::AddFunction(const GlobalVar& gvar, const PrimFunc& f) { - func_has_pstate_sm = false; - func_has_pstate_za = false; + llvm::Function* func = DeclareFunction(gvar, f); + current_prim_func_ = f; + packed_function_ = f->GetAttr(tvm::attr::kCallingConv) == CallingConv::kCPackedFunc + ? func + : nullptr; CodeGenCPU::AddFunction(gvar, f); + current_prim_func_ = std::nullopt; + packed_function_ = nullptr; +} + +void CodeGenAArch64::SetComputeScopeAttributes(llvm::Function* func) { + CodeGenCPU::SetComputeScopeAttributes(func); + // MakePackedAPI keeps PrimFunc attrs on the packed wrapper, but the SME contract belongs + // to its outlined compute function. Runtime callbacks and nested helpers have separate ABIs. + if (function_ == packed_function_) { + SetPStateAttributes(func, current_prim_func_.value()); + } +} + +void CodeGenAArch64::SetPStateAttributes(llvm::Function* func, const PrimFunc& f) { + // These string PrimFunc attrs are exposed directly through T.func_attr and with_attr. + if (auto sm = f->GetAttr("aarch64_pstate_sm")) { + // A locally streaming body does not change a bodyless declaration's interface. + if (sm.value() != "body" || f->body.has_value()) { + func->addFnAttr(MakeStringRef("aarch64_pstate_sm_" + sm.value())); + } + } + if (auto za = f->GetAttr("aarch64_pstate_za")) { +#if TVM_LLVM_VERSION >= 190 + // LLVM 19 renamed the new/shared policies. Keep the legacy spelling for preserved: + // aarch64_preserves_za describes a shared interface, unlike the old private interface. + if (za.value() == "new") { + func->addFnAttr("aarch64_new_za"); + } else if (za.value() == "shared") { + func->addFnAttr("aarch64_inout_za"); + } else +#endif + { + func->addFnAttr(MakeStringRef("aarch64_pstate_za_" + za.value())); + } + } } void CodeGenAArch64::SetTargetAttributes(llvm::Function* func) { @@ -80,46 +130,6 @@ void CodeGenAArch64::SetTargetAttributes(llvm::Function* func) { CodeGenCPU::SetTargetAttributes(func); } -/*! - * \brief Visit and handle AArch64 specific pragmas. To be AArch64 specific, - * the expectation is that they are prepended with "pragma_aarch64". - */ -void CodeGenAArch64::Dispatch_(const AttrStmtNode* op) { - std::string attr_key = op->attr_key; - - if (!tirx::attr::IsPragmaKey(attr_key)) { - CodeGenCPU::Dispatch_(op); - return; - } - bool is_aarch64_specific_pragma = attr_key.substr(7, 7) == "aarch64"; - if (!is_aarch64_specific_pragma) { - CodeGenCPU::Dispatch_(op); - return; - } - - const auto* attr_value = op->value.as(); - TVM_FFI_ICHECK(attr_value) << "Expect " << attr_key << " to have a ffi::String value but was " - << op->value->GetTypeKey(); - - std::string aarch64_attr_key = attr_key.substr(7); - if (aarch64_attr_key == "aarch64_pstate_sm") { - TVM_FFI_ICHECK(!func_has_pstate_sm) - << "Multiple definitions of " << op->attr_key << " attribute found in the function " - << function_->getName().data(); - function_->addFnAttr({aarch64_attr_key + "_" + attr_value->value}); - func_has_pstate_sm = true; - } else if (aarch64_attr_key == "aarch64_pstate_za") { - TVM_FFI_ICHECK(!func_has_pstate_za) - << "Multiple definitions of " << op->attr_key << " attribute found in the function " - << function_->getName().data(); - function_->addFnAttr({aarch64_attr_key + "_" + attr_value->value}); - func_has_pstate_za = true; - } else { - LOG(WARNING) << "Unknown pragma " << op->attr_key; - } - this->Dispatch(op->body); -} - TVM_FFI_STATIC_INIT_BLOCK() { namespace refl = tvm::ffi::reflection; refl::GlobalDef().def_packed("tvm.codegen.llvm.target_aarch64", diff --git a/src/target/llvm/codegen_cpu.cc b/src/target/llvm/codegen_cpu.cc index 7833a3221556..20a287943b62 100644 --- a/src/target/llvm/codegen_cpu.cc +++ b/src/target/llvm/codegen_cpu.cc @@ -559,7 +559,7 @@ void CodeGenCPU::CreateComputeScope(const AttrStmtNode* op) { TVM_FFI_ICHECK(value != nullptr); llvm::Function* fcompute = llvm::Function::Create(ftype, llvm::Function::InternalLinkage, MakeStringRef(value->value), module_.get()); - SetTargetAttributes(fcompute); + SetComputeScopeAttributes(fcompute); for (auto it = fcompute->arg_begin(); it != fcompute->arg_end(); it++) { const Var& var = vargs[std::distance(fcompute->arg_begin(), it)]; it->setName(std::string(var->name)); diff --git a/src/target/llvm/codegen_cpu.h b/src/target/llvm/codegen_cpu.h index dd9dc57436a2..a73210d151bf 100644 --- a/src/target/llvm/codegen_cpu.h +++ b/src/target/llvm/codegen_cpu.h @@ -79,6 +79,8 @@ class CodeGenCPU : public CodeGenLLVM { const ffi::Array& args, bool skip_first_arg) override; protected: + // Set attributes on the function outlined from a compute scope. + virtual void SetComputeScopeAttributes(llvm::Function* func) { SetTargetAttributes(func); } void AddStartupFunction() final; // meta data llvm::MDNode* md_tbaa_ctx_ptr_{nullptr};