clang-hipify: add Replacement Excludes

Excludes are not replaced, for instance, CHECK_CUDA_ERROR and CUDA_SAFE_CALL.
Add check for excludes in MacroExpands and CallExpr routines.
This commit is contained in:
Evgeny Mankov
2016-07-01 19:58:14 +03:00
parent 2eac7144f0
commit fd1e556cf2
+66 -48
View File
@@ -46,6 +46,7 @@ THE SOFTWARE.
#include <cstdio> #include <cstdio>
#include <fstream> #include <fstream>
#include <set>
using namespace clang; using namespace clang;
using namespace clang::ast_matchers; using namespace clang::ast_matchers;
@@ -82,6 +83,10 @@ namespace {
struct cuda2hipMap { struct cuda2hipMap {
cuda2hipMap() { cuda2hipMap() {
// Replacement Excludes
cudaExcludes = {"CHECK_CUDA_ERROR", "CUDA_SAFE_CALL"};
// Defines // Defines
cuda2hipRename["__CUDACC__"] = {"__HIPCC__", CONV_DEF}; cuda2hipRename["__CUDACC__"] = {"__HIPCC__", CONV_DEF};
@@ -89,6 +94,9 @@ struct cuda2hipMap {
cuda2hipRename["cuda_runtime.h"] = {"hip_runtime.h", CONV_INCLUDE}; cuda2hipRename["cuda_runtime.h"] = {"hip_runtime.h", CONV_INCLUDE};
cuda2hipRename["cuda_runtime_api.h"] = {"hip_runtime_api.h", CONV_INCLUDE}; cuda2hipRename["cuda_runtime_api.h"] = {"hip_runtime_api.h", CONV_INCLUDE};
// HIP includes
cuda2hipRename["cudacommon.h.prehip"] = {"cudacommon.h", CONV_INCLUDE};
// CUBLAS includes // CUBLAS includes
cuda2hipRename["cublas.h"] = {"hipblas.h", CONV_INCLUDE}; cuda2hipRename["cublas.h"] = {"hipblas.h", CONV_INCLUDE};
cuda2hipRename["cublas_v2.h"] = {"hipblas.h", CONV_INCLUDE}; cuda2hipRename["cublas_v2.h"] = {"hipblas.h", CONV_INCLUDE};
@@ -941,7 +949,7 @@ struct cuda2hipMap {
// ROTMG // ROTMG
//cuda2hipRename["cublasSrotmg_v2"] = {"hipblasSrotmg", CONV_BLAS}; //cuda2hipRename["cublasSrotmg_v2"] = {"hipblasSrotmg", CONV_BLAS};
//cuda2hipRename["cublasDrotmg_v2"] = {"hipblasDrotmg", CONV_BLAS}; //cuda2hipRename["cublasDrotmg_v2"] = {"hipblasDrotmg", CONV_BLAS};
} }
struct HipNames { struct HipNames {
StringRef hipName; StringRef hipName;
@@ -949,6 +957,7 @@ struct cuda2hipMap {
}; };
SmallDenseMap<StringRef, HipNames> cuda2hipRename; SmallDenseMap<StringRef, HipNames> cuda2hipRename;
std::set<StringRef> cudaExcludes;
}; };
StringRef unquoteStr(StringRef s) { StringRef unquoteStr(StringRef s) {
@@ -1055,47 +1064,48 @@ struct HipifyPPCallbacks : public PPCallbacks, public SourceFileCallbacks {
const MacroDefinition &MD, SourceRange Range, const MacroDefinition &MD, SourceRange Range,
const MacroArgs *Args) override { const MacroArgs *Args) override {
if (_sm->isWrittenInMainFile(MacroNameTok.getLocation())) { if (_sm->isWrittenInMainFile(MacroNameTok.getLocation())) {
for (unsigned int i = 0; Args && i < MD.getMacroInfo()->getNumArgs(); StringRef macroName = MacroNameTok.getIdentifierInfo()->getName();
i++) { if (N.cudaExcludes.end() == N.cudaExcludes.find(macroName)) {
StringRef macroName = MacroNameTok.getIdentifierInfo()->getName(); for (unsigned int i = 0; Args && i < MD.getMacroInfo()->getNumArgs(); i++) {
std::vector<Token> toks; std::vector<Token> toks;
// Code below is a kind of stolen from 'MacroArgs::getPreExpArgument' // Code below is a kind of stolen from 'MacroArgs::getPreExpArgument'
// to workaround the 'const' MacroArgs passed into this hook. // to workaround the 'const' MacroArgs passed into this hook.
const Token *start = Args->getUnexpArgument(i); const Token *start = Args->getUnexpArgument(i);
size_t len = Args->getArgLength(start) + 1; size_t len = Args->getArgLength(start) + 1;
#if (LLVM_VERSION_MAJOR >= 3) && (LLVM_VERSION_MINOR >= 9) #if (LLVM_VERSION_MAJOR >= 3) && (LLVM_VERSION_MINOR >= 9)
_pp->EnterTokenStream(ArrayRef<Token>(start, len), false); _pp->EnterTokenStream(ArrayRef<Token>(start, len), false);
#else #else
_pp->EnterTokenStream(start, len, false, false); _pp->EnterTokenStream(start, len, false, false);
#endif #endif
do { do {
toks.push_back(Token()); toks.push_back(Token());
Token &tk = toks.back(); Token &tk = toks.back();
_pp->Lex(tk); _pp->Lex(tk);
} while (toks.back().isNot(tok::eof)); } while (toks.back().isNot(tok::eof));
_pp->RemoveTopOfLexerStack(); _pp->RemoveTopOfLexerStack();
// end of stolen code // end of stolen code
for (auto tok : toks) { for (auto tok : toks) {
if (tok.isAnyIdentifier()) { if (tok.isAnyIdentifier()) {
StringRef name = tok.getIdentifierInfo()->getName(); StringRef name = tok.getIdentifierInfo()->getName();
const auto found = N.cuda2hipRename.find(name); const auto found = N.cuda2hipRename.find(name);
if (found != N.cuda2hipRename.end()) { if (found != N.cuda2hipRename.end()) {
countReps[found->second.countType]++; countReps[found->second.countType]++;
StringRef repName = found->second.hipName; StringRef repName = found->second.hipName;
DEBUG(dbgs() DEBUG(dbgs()
<< "Identifier " << name << "Identifier " << name
<< " found as an actual argument in expansion of macro " << " found as an actual argument in expansion of macro "
<< macroName << "\n" << macroName << "\n"
<< "will be replaced with: " << repName << "\n"); << "will be replaced with: " << repName << "\n");
SourceLocation sl = tok.getLocation(); SourceLocation sl = tok.getLocation();
Replacement Rep(*_sm, sl, name.size(), repName); Replacement Rep(*_sm, sl, name.size(), repName);
Replace->insert(Rep); Replace->insert(Rep);
}
}
if (tok.is(tok::string_literal)) {
StringRef s(tok.getLiteralData(), tok.getLength());
processString(unquoteStr(s), N, Replace, *_sm, tok.getLocation(),
countReps);
} }
}
if (tok.is(tok::string_literal)) {
StringRef s(tok.getLiteralData(), tok.getLength());
processString(unquoteStr(s), N, Replace, *_sm, tok.getLocation(),
countReps);
} }
} }
} }
@@ -1168,21 +1178,29 @@ public:
StringRef name = funcDcl->getDeclName().getAsString(); StringRef name = funcDcl->getDeclName().getAsString();
const auto found = N.cuda2hipRename.find(name); const auto found = N.cuda2hipRename.find(name);
if (found != N.cuda2hipRename.end()) { if (found != N.cuda2hipRename.end()) {
countReps[found->second.countType]++;
StringRef repName = found->second.hipName; StringRef repName = found->second.hipName;
SourceLocation sl = call->getLocStart(); SourceLocation sl = call->getLocStart();
size_t length = name.size(); size_t length = name.size();
bool bReplace = true;
if (SM->isMacroArgExpansion(sl)) { if (SM->isMacroArgExpansion(sl)) {
sl = SM->getImmediateSpellingLoc(sl); sl = SM->getImmediateSpellingLoc(sl);
} }
else if (SM->isMacroBodyExpansion(sl)) { else if (SM->isMacroBodyExpansion(sl)) {
sl = SM->getExpansionLoc(sl); SourceLocation sl_macro = SM->getExpansionLoc(sl);
SourceLocation sl_end = SourceLocation sl_end = Lexer::getLocForEndOfToken(sl_macro, 0, *SM, DefaultLangOptions);
Lexer::getLocForEndOfToken(sl, 0, *SM, DefaultLangOptions); length = SM->getCharacterData(sl_end) - SM->getCharacterData(sl_macro);
length = SM->getCharacterData(sl_end) - SM->getCharacterData(sl); StringRef macroName = StringRef(SM->getCharacterData(sl_macro), length);
if (N.cudaExcludes.end() != N.cudaExcludes.find(macroName)) {
bReplace = false;
} else {
sl = sl_macro;
}
}
if (bReplace) {
countReps[found->second.countType]++;
Replacement Rep(*SM, sl, length, repName);
Replace->insert(Rep);
} }
Replacement Rep(*SM, sl, length, repName);
Replace->insert(Rep);
} }
} }