#include "mlir/Dialect/LLVMIR/BasicPtxBuilderInterface.h"
#include "mlir/IR/BuiltinTypes.h"
#include "mlir/IR/Diagnostics.h"
#include "mlir/IR/Location.h"
#include "mlir/IR/MLIRContext.h"
#include "mlir/Support/LLVM.h"
#include "llvm/ADT/StringExtras.h"
#include "llvm/ADT/TypeSwitch.h"
#include "llvm/Support/DebugLog.h"
#include "llvm/Support/FormatVariadic.h"
#include "llvm/Support/LogicalResult.h"
#include "llvm/Support/Regex.h"
#define DEBUG_TYPE "ptx-builder"
#include "mlir/Dialect/LLVMIR/BasicPtxBuilderInterface.cpp.inc"
using namespace mlir;
using namespace NVVM;
static constexpr int64_t kSharedMemorySpace = 3;
static FailureOr<char> getRegisterType(Type type, Location loc) {
MLIRContext *ctx = type.getContext();
auto i16 = IntegerType::get(ctx, 16);
auto i32 = IntegerType::get(ctx, 32);
auto f32 = Float32Type::get(ctx);
auto getRegisterTypeForScalar = [&](Type type) -> FailureOr<char> {
if (type.isInteger(1))
return 'b';
if (type.isInteger(16))
return 'h';
if (type.isInteger(32))
return 'r';
if (type.isInteger(64))
return 'l';
if (type.isF32())
return 'f';
if (type.isF64())
return 'd';
if (auto ptr = dyn_cast<LLVM::LLVMPointerType>(type)) {
if (ptr.getAddressSpace() == kSharedMemorySpace) {
return 'r';
}
return 'l';
}
mlir::emitError(
loc, "The register type could not be deduced from MLIR type. The ")
<< type
<< " is not supported. Supported types are:"
"i1, i16, i32, i64, f32, f64,"
"pointers.\nPlease use llvm.bitcast if you have different type. "
"\nSee the constraints from here: "
"https://docs.nvidia.com/cuda/inline-ptx-assembly/"
"index.html#constraints";
return failure();
};
if (auto v = dyn_cast<VectorType>(type)) {
assert(v.getNumDynamicDims() == 0 && "Dynamic vectors are not supported");
int64_t lanes = v.getNumElements();
Type elem = v.getElementType();
if (lanes <= 1)
return getRegisterTypeForScalar(elem);
Type widened = elem;
switch (lanes) {
case 2:
if (elem.isF16() || elem.isBF16())
widened = f32;
else if (elem.isFloat(8))
widened = i16;
break;
case 4:
if (elem.isInteger(8))
widened = i32;
else if (elem.isFloat(8))
widened = f32;
else if (elem.isFloat(4))
widened = i16;
break;
default:
break;
}
return getRegisterTypeForScalar(widened);
}
return getRegisterTypeForScalar(type);
}
static FailureOr<char> getRegisterType(Value v, Location loc) {
if (v.getDefiningOp<LLVM::ConstantOp>())
return 'n';
return getRegisterType(v.getType(), loc);
}
static SmallVector<Value> extractStructElements(PatternRewriter &rewriter,
Location loc, Value structVal) {
auto structTy = dyn_cast<LLVM::LLVMStructType>(structVal.getType());
assert(structTy && "expected LLVM struct");
SmallVector<Value> elems;
for (unsigned i : llvm::seq<unsigned>(0, structTy.getBody().size()))
elems.push_back(LLVM::ExtractValueOp::create(rewriter, loc, structVal, i));
return elems;
}
LogicalResult PtxBuilder::insertValue(Value v, PTXRegisterMod itype) {
LDBG() << v << "\t Modifier : " << itype << "\n";
registerModifiers.push_back(itype);
Location loc = interfaceOp->getLoc();
auto getModifier = [&]() -> const char * {
switch (itype) {
case PTXRegisterMod::Read:
return "";
case PTXRegisterMod::Write:
return "=";
case PTXRegisterMod::ReadWrite:
return "+";
}
llvm_unreachable("Unknown PTX register modifier");
};
auto addValue = [&](Value v) {
if (itype == PTXRegisterMod::Read) {
ptxOperands.push_back(v);
return;
}
if (itype == PTXRegisterMod::ReadWrite)
ptxOperands.push_back(v);
hasResult = true;
};
llvm::raw_string_ostream ss(registerConstraints);
if (auto stype = dyn_cast<LLVM::LLVMStructType>(v.getType())) {
if (itype == PTXRegisterMod::Write) {
addValue(v);
}
for (auto [idx, t] : llvm::enumerate(stype.getBody())) {
if (itype != PTXRegisterMod::Write) {
Value extractValue =
LLVM::ExtractValueOp::create(rewriter, loc, v, idx);
addValue(extractValue);
}
if (itype == PTXRegisterMod::ReadWrite) {
ss << idx << ",";
} else {
FailureOr<char> regType = getRegisterType(t, loc);
if (failed(regType))
return rewriter.notifyMatchFailure(loc,
"failed to get register type");
ss << getModifier() << regType.value() << ",";
}
}
return success();
}
addValue(v);
FailureOr<char> regType = getRegisterType(v, loc);
if (failed(regType))
return rewriter.notifyMatchFailure(loc, "failed to get register type");
ss << getModifier() << regType.value() << ",";
return success();
}
static bool
needsPackUnpack(BasicPtxBuilderInterface interfaceOp,
bool needsManualRegisterMapping,
SmallVectorImpl<PTXRegisterMod> ®isterModifiers) {
if (needsManualRegisterMapping)
return false;
const unsigned writeOnlyVals = interfaceOp->getNumResults();
const unsigned readWriteVals =
llvm::count_if(registerModifiers, [](PTXRegisterMod m) {
return m == PTXRegisterMod::ReadWrite;
});
return (writeOnlyVals + readWriteVals) > 1;
}
static SmallVector<Type>
packResultTypes(BasicPtxBuilderInterface interfaceOp,
bool needsManualRegisterMapping,
SmallVectorImpl<PTXRegisterMod> ®isterModifiers,
SmallVectorImpl<Value> &ptxOperands) {
MLIRContext *ctx = interfaceOp->getContext();
TypeRange resultRange = interfaceOp->getResultTypes();
if (!needsPackUnpack(interfaceOp, needsManualRegisterMapping,
registerModifiers)) {
if (interfaceOp->getResults().size() == 1)
return SmallVector<Type>{resultRange.front()};
for (auto [m, v] : llvm::zip(registerModifiers, ptxOperands))
if (m == PTXRegisterMod::ReadWrite)
return SmallVector<Type>{v.getType()};
}
SmallVector<Type> packed;
for (auto [m, v] : llvm::zip(registerModifiers, ptxOperands))
if (m == PTXRegisterMod::ReadWrite)
packed.push_back(v.getType());
for (Type t : resultRange)
packed.push_back(t);
if (packed.empty())
return {};
auto sTy = LLVM::LLVMStructType::getLiteral(ctx, packed, false);
return SmallVector<Type>{sTy};
}
static std::string canonicalizeRegisterConstraints(llvm::StringRef csv) {
SmallVector<llvm::StringRef> toks;
SmallVector<std::string> out;
SmallVector<unsigned> plusIdx;
csv.split(toks, ',');
out.reserve(toks.size() + 8);
for (unsigned i = 0, e = toks.size(); i < e; ++i) {
StringRef t = toks[i].trim();
if (t.consume_front("+")) {
plusIdx.push_back(i);
out.push_back(("=" + t).str());
} else {
out.push_back(t.str());
}
}
for (unsigned idx : plusIdx)
out.push_back(std::to_string(idx));
std::string result;
result.reserve(csv.size() + plusIdx.size() * 2);
llvm::raw_string_ostream os(result);
for (size_t i = 0; i < out.size(); ++i) {
if (i)
os << ',';
os << out[i];
}
return os.str();
}
constexpr llvm::StringLiteral kReadWritePrefix{"rw"};
constexpr llvm::StringLiteral kWriteOnlyPrefix{"w"};
constexpr llvm::StringLiteral kReadOnlyPrefix{"r"};
static llvm::Regex getPredicateMappingRegex() {
llvm::Regex rx(llvm::formatv(R"(\{\$({0}|{1}|{2})([0-9]+)\})",
kReadWritePrefix, kWriteOnlyPrefix,
kReadOnlyPrefix)
.str());
return rx;
}
void mlir::NVVM::countPlaceholderNumbers(
StringRef ptxCode, llvm::SmallDenseSet<unsigned int> &seenRW,
llvm::SmallDenseSet<unsigned int> &seenW,
llvm::SmallDenseSet<unsigned int> &seenR,
llvm::SmallVectorImpl<unsigned int> &rwNums,
llvm::SmallVectorImpl<unsigned int> &wNums,
llvm::SmallVectorImpl<unsigned int> &rNums) {
llvm::Regex rx = getPredicateMappingRegex();
StringRef rest = ptxCode;
SmallVector<StringRef, 3> m;
while (!rest.empty() && rx.match(rest, &m)) {
unsigned num = 0;
(void)m[2].getAsInteger(10, num);
if (m[1].equals_insensitive(kReadWritePrefix)) {
if (seenRW.insert(num).second)
rwNums.push_back(num);
} else if (m[1].equals_insensitive(kWriteOnlyPrefix)) {
if (seenW.insert(num).second)
wNums.push_back(num);
} else {
if (seenR.insert(num).second)
rNums.push_back(num);
}
const size_t advance = (size_t)(m[0].data() - rest.data()) + m[0].size();
rest = rest.drop_front(advance);
}
}
static std::string rewriteAsmPlaceholders(llvm::StringRef ptxCode) {
llvm::SmallDenseSet<unsigned> seenRW, seenW, seenR;
llvm::SmallVector<unsigned> rwNums, wNums, rNums;
countPlaceholderNumbers(ptxCode, seenRW, seenW, seenR, rwNums, wNums, rNums);
llvm::sort(rwNums);
llvm::sort(wNums);
llvm::sort(rNums);
llvm::DenseMap<unsigned, unsigned> rwMap, wMap, rMap;
unsigned nextId = 0;
for (unsigned n : rwNums)
rwMap[n] = nextId++;
for (unsigned n : wNums)
wMap[n] = nextId++;
for (unsigned n : rNums)
rMap[n] = nextId++;
std::string out;
out.reserve(ptxCode.size());
size_t prev = 0;
StringRef rest = ptxCode;
SmallVector<StringRef, 3> matches;
llvm::Regex rx = getPredicateMappingRegex();
while (!rest.empty() && rx.match(rest, &matches)) {
size_t absStart = (size_t)(matches[0].data() - ptxCode.data());
size_t absEnd = absStart + matches[0].size();
out.append(ptxCode.data() + prev, ptxCode.data() + absStart);
unsigned num = 0;
(void)matches[2].getAsInteger(10, num);
unsigned id = 0;
if (matches[1].equals_insensitive(kReadWritePrefix))
id = rwMap.lookup(num);
else if (matches[1].equals_insensitive(kWriteOnlyPrefix))
id = wMap.lookup(num);
else
id = rMap.lookup(num);
out.push_back('$');
out += std::to_string(id);
prev = absEnd;
const size_t advance =
(size_t)(matches[0].data() - rest.data()) + matches[0].size();
rest = rest.drop_front(advance);
}
out.append(ptxCode.data() + prev, ptxCode.data() + ptxCode.size());
return out;
}
LLVM::InlineAsmOp PtxBuilder::build() {
auto asmDialectAttr = LLVM::AsmDialectAttr::get(interfaceOp->getContext(),
LLVM::AsmDialect::AD_ATT);
SmallVector<Type> resultTypes = packResultTypes(
interfaceOp, needsManualRegisterMapping, registerModifiers, ptxOperands);
if (!registerConstraints.empty() &&
registerConstraints[registerConstraints.size() - 1] == ',')
registerConstraints.pop_back();
registerConstraints = canonicalizeRegisterConstraints(registerConstraints);
std::string ptxInstruction = interfaceOp.getPtx();
if (!needsManualRegisterMapping)
ptxInstruction = rewriteAsmPlaceholders(ptxInstruction);
if (interfaceOp.getPredicate().has_value() &&
interfaceOp.getPredicate().value()) {
std::string predicateStr = "@%";
predicateStr += std::to_string((ptxOperands.size() - 1));
ptxInstruction = predicateStr + " " + ptxInstruction;
}
llvm::replace(ptxInstruction, '%', '$');
return LLVM::InlineAsmOp::create(
rewriter, interfaceOp->getLoc(),
resultTypes,
ptxOperands,
ptxInstruction,
registerConstraints.data(),
interfaceOp.hasSideEffect(),
false, LLVM::TailCallKind::None,
asmDialectAttr,
ArrayAttr());
}
void PtxBuilder::buildAndReplaceOp() {
LLVM::InlineAsmOp inlineAsmOp = build();
LDBG() << "\n Generated PTX \n\t" << inlineAsmOp;
if (!hasResult) {
rewriter.eraseOp(interfaceOp);
return;
}
if (needsManualRegisterMapping) {
rewriter.replaceOp(interfaceOp, inlineAsmOp->getResults());
return;
}
if (!needsPackUnpack(interfaceOp, needsManualRegisterMapping,
registerModifiers)) {
if (inlineAsmOp->getNumResults() > 0) {
rewriter.replaceOp(interfaceOp, inlineAsmOp->getResults());
} else {
SmallVector<Value> results;
for (auto [m, v] : llvm::zip(registerModifiers, ptxOperands))
if (m == PTXRegisterMod::ReadWrite) {
results.push_back(v);
break;
}
rewriter.replaceOp(interfaceOp, results);
}
return;
}
const bool hasRW = llvm::any_of(registerModifiers, [](PTXRegisterMod m) {
return m == PTXRegisterMod::ReadWrite;
});
assert(LLVM::LLVMStructType::classof(inlineAsmOp.getResultTypes().front()) &&
"expected struct return for multi-result inline asm");
Value structVal = inlineAsmOp.getResult(0);
SmallVector<Value> unpacked =
extractStructElements(rewriter, interfaceOp->getLoc(), structVal);
if (!hasRW && interfaceOp->getResults().size() > 0) {
rewriter.replaceOp(interfaceOp, unpacked);
return;
}
if (hasRW && interfaceOp->getResults().size() == 0) {
unsigned idx = 0;
for (auto [m, v] : llvm::zip(registerModifiers, ptxOperands)) {
if (m != PTXRegisterMod::ReadWrite)
continue;
Value repl = unpacked[idx++];
v.replaceUsesWithIf(repl, [&](OpOperand &use) {
Operation *owner = use.getOwner();
return owner != interfaceOp && owner != inlineAsmOp;
});
}
rewriter.eraseOp(interfaceOp);
return;
}
{
unsigned idx = 0;
for (auto [m, v] : llvm::zip(registerModifiers, ptxOperands)) {
if (m != PTXRegisterMod::ReadWrite)
continue;
Value repl = unpacked[idx++];
v.replaceUsesWithIf(repl, [&](OpOperand &use) {
Operation *owner = use.getOwner();
return owner != interfaceOp && owner != inlineAsmOp;
});
}
SmallVector<Value> tail;
tail.reserve(unpacked.size() - idx);
for (unsigned i = idx, e = unpacked.size(); i < e; ++i)
tail.push_back(unpacked[i]);
rewriter.replaceOp(interfaceOp, tail);
}
}