[HIPIFY][CUB][#1460] Implement cubFunctionTemplateDecl matcher
+ Add cub_02.cu test + Partial fixes #1460
This commit is contained in:
@@ -61,6 +61,7 @@ const StringRef sCudaLaunchKernel = "cudaLaunchKernel";
|
||||
const StringRef sCudaHostFuncCall = "cudaHostFuncCall";
|
||||
const StringRef sCudaDeviceFuncCall = "cudaDeviceFuncCall";
|
||||
const StringRef sCubNamespacePrefix = "cubNamespacePrefix";
|
||||
const StringRef sCubFunctionTemplateDecl = "cubFunctionTemplateDecl";
|
||||
|
||||
std::set<std::string> DeviceSymbolFunctions0 {
|
||||
{sCudaMemcpyToSymbol},
|
||||
@@ -449,6 +450,41 @@ bool HipifyAction::cubNamespacePrefix(const mat::MatchFinder::MatchResult &Resul
|
||||
return false;
|
||||
}
|
||||
|
||||
bool HipifyAction::cubFunctionTemplateDecl(const mat::MatchFinder::MatchResult &Result) {
|
||||
if (auto *decl = Result.Nodes.getNodeAs<clang::FunctionTemplateDecl>(sCubFunctionTemplateDecl)) {
|
||||
auto *Tparams = decl->getTemplateParameters();
|
||||
bool ret = false;
|
||||
for (size_t I = 0; I < Tparams->size(); ++I) {
|
||||
const clang::ValueDecl *valueDecl = dyn_cast<clang::ValueDecl>(Tparams->getParam(I));
|
||||
if (!valueDecl) continue;
|
||||
clang::QualType QT = valueDecl->getType();
|
||||
auto *t = QT.getTypePtr();
|
||||
if (!t) continue;
|
||||
const clang::ElaboratedType *et = t->getAs<clang::ElaboratedType>();
|
||||
if (!et) continue;
|
||||
const clang::NestedNameSpecifier *nns = et->getQualifier();
|
||||
if (!nns) continue;
|
||||
const clang::NamespaceDecl *nsd = nns->getAsNamespace();
|
||||
if (!nsd) continue;
|
||||
const clang::SourceRange sr = valueDecl->getSourceRange();
|
||||
clang::SourceLocation sl(sr.getBegin());
|
||||
clang::SourceLocation end(sr.getEnd());
|
||||
auto &SM = getCompilerInstance().getSourceManager();
|
||||
size_t length = SM.getCharacterData(end) - SM.getCharacterData(sl);
|
||||
StringRef sfull = StringRef(SM.getCharacterData(sl), length);
|
||||
std::string name = nsd->getDeclName().getAsString();
|
||||
size_t offset = sfull.find(name);
|
||||
if (offset > 0) {
|
||||
sl = sl.getLocWithOffset(offset);
|
||||
}
|
||||
FindAndReplace(name, sl, CUDA_CUB_TYPE_NAME_MAP);
|
||||
ret = true;
|
||||
}
|
||||
return ret;
|
||||
}
|
||||
return false;
|
||||
}
|
||||
|
||||
bool HipifyAction::cudaHostFuncCall(const mat::MatchFinder::MatchResult &Result) {
|
||||
if (auto *call = Result.Nodes.getNodeAs<clang::CallExpr>(sCudaHostFuncCall)) {
|
||||
if (!call->getNumArgs()) return false;
|
||||
@@ -555,6 +591,13 @@ std::unique_ptr<clang::ASTConsumer> HipifyAction::CreateASTConsumer(clang::Compi
|
||||
).bind(sCubNamespacePrefix),
|
||||
this
|
||||
);
|
||||
// TODO: Maybe worth to make it more concrete based on final cubFunctionTemplateDecl
|
||||
Finder->addMatcher(
|
||||
mat::functionTemplateDecl(
|
||||
mat::isExpansionInMainFile()
|
||||
).bind(sCubFunctionTemplateDecl),
|
||||
this
|
||||
);
|
||||
// Ownership is transferred to the caller.
|
||||
return Finder->newASTConsumer();
|
||||
}
|
||||
@@ -668,4 +711,5 @@ void HipifyAction::run(const mat::MatchFinder::MatchResult &Result) {
|
||||
if (cudaHostFuncCall(Result)) return;
|
||||
if (cudaDeviceFuncCall(Result)) return;
|
||||
if (cubNamespacePrefix(Result)) return;
|
||||
if (cubFunctionTemplateDecl(Result)) return;
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user