Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
104 changes: 57 additions & 47 deletions src/target/llvm/codegen_aarch64.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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<PrimFunc> 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<CallingConv>(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<CallingConv>(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<ffi::String>("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<ffi::String>("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) {
Expand Down Expand Up @@ -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<StringImmNode>();
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",
Expand Down
2 changes: 1 addition & 1 deletion src/target/llvm/codegen_cpu.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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));
Expand Down
2 changes: 2 additions & 0 deletions src/target/llvm/codegen_cpu.h
Original file line number Diff line number Diff line change
Expand Up @@ -79,6 +79,8 @@ class CodeGenCPU : public CodeGenLLVM {
const ffi::Array<Expr>& 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};
Expand Down
Loading