542 lines · cpp
1//===- OpenACCToLLVMIRTranslation.cpp -------------------------------------===//2//3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.4// See https://llvm.org/LICENSE.txt for license information.5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception6//7//===----------------------------------------------------------------------===//8//9// This file implements a translation between the MLIR OpenACC dialect and LLVM10// IR.11//12//===----------------------------------------------------------------------===//13 14#include "mlir/Target/LLVMIR/Dialect/OpenACC/OpenACCToLLVMIRTranslation.h"15#include "mlir/Analysis/TopologicalSortUtils.h"16#include "mlir/Dialect/LLVMIR/LLVMDialect.h"17#include "mlir/Dialect/OpenACC/OpenACC.h"18#include "mlir/IR/Operation.h"19#include "mlir/Support/LLVM.h"20#include "mlir/Target/LLVMIR/Dialect/OpenMPCommon.h"21#include "mlir/Target/LLVMIR/ModuleTranslation.h"22 23#include "llvm/ADT/TypeSwitch.h"24#include "llvm/Frontend/OpenMP/OMPConstants.h"25 26using namespace mlir;27 28using OpenACCIRBuilder = llvm::OpenMPIRBuilder;29 30//===----------------------------------------------------------------------===//31// Utility functions32//===----------------------------------------------------------------------===//33 34/// Flag values are extracted from openmp/libomptarget/include/omptarget.h and35/// mapped to corresponding OpenACC flags.36static constexpr uint64_t kCreateFlag = 0x000;37static constexpr uint64_t kDeviceCopyinFlag = 0x001;38static constexpr uint64_t kHostCopyoutFlag = 0x002;39static constexpr uint64_t kPresentFlag = 0x1000;40static constexpr uint64_t kDeleteFlag = 0x008;41// Runtime extension to implement the OpenACC second reference counter.42static constexpr uint64_t kHoldFlag = 0x2000;43 44/// Default value for the device id45static constexpr int64_t kDefaultDevice = -1;46 47/// Create the location struct from the operation location information.48static llvm::Value *createSourceLocationInfo(OpenACCIRBuilder &builder,49 Operation *op) {50 auto loc = op->getLoc();51 auto funcOp = op->getParentOfType<LLVM::LLVMFuncOp>();52 StringRef funcName = funcOp ? funcOp.getName() : "unknown";53 uint32_t strLen;54 llvm::Constant *locStr = mlir::LLVM::createSourceLocStrFromLocation(55 loc, builder, funcName, strLen);56 return builder.getOrCreateIdent(locStr, strLen);57}58 59/// Return the runtime function used to lower the given operation.60static llvm::Function *getAssociatedFunction(OpenACCIRBuilder &builder,61 Operation *op) {62 return llvm::TypeSwitch<Operation *, llvm::Function *>(op)63 .Case([&](acc::EnterDataOp) {64 return builder.getOrCreateRuntimeFunctionPtr(65 llvm::omp::OMPRTL___tgt_target_data_begin_mapper);66 })67 .Case([&](acc::ExitDataOp) {68 return builder.getOrCreateRuntimeFunctionPtr(69 llvm::omp::OMPRTL___tgt_target_data_end_mapper);70 })71 .Case([&](acc::UpdateOp) {72 return builder.getOrCreateRuntimeFunctionPtr(73 llvm::omp::OMPRTL___tgt_target_data_update_mapper);74 });75 llvm_unreachable("Unknown OpenACC operation");76}77 78/// Extract pointer, size and mapping information from operands79/// to populate the future functions arguments.80static LogicalResult81processOperands(llvm::IRBuilderBase &builder,82 LLVM::ModuleTranslation &moduleTranslation, Operation *op,83 ValueRange operands, unsigned totalNbOperand,84 uint64_t operandFlag, SmallVector<uint64_t> &flags,85 SmallVectorImpl<llvm::Constant *> &names, unsigned &index,86 struct OpenACCIRBuilder::MapperAllocas &mapperAllocas) {87 OpenACCIRBuilder *accBuilder = moduleTranslation.getOpenMPBuilder();88 llvm::LLVMContext &ctx = builder.getContext();89 auto *i8PtrTy = llvm::PointerType::getUnqual(ctx);90 auto *arrI8PtrTy = llvm::ArrayType::get(i8PtrTy, totalNbOperand);91 auto *i64Ty = llvm::Type::getInt64Ty(ctx);92 auto *arrI64Ty = llvm::ArrayType::get(i64Ty, totalNbOperand);93 94 for (Value data : operands) {95 llvm::Value *dataValue = moduleTranslation.lookupValue(data);96 97 llvm::Value *dataPtrBase;98 llvm::Value *dataPtr;99 llvm::Value *dataSize;100 101 if (isa<LLVM::LLVMPointerType>(data.getType())) {102 dataPtrBase = dataValue;103 dataPtr = dataValue;104 dataSize = accBuilder->getSizeInBytes(dataValue);105 } else {106 return op->emitOpError()107 << "Data operand must be legalized before translation."108 << "Unsupported type: " << data.getType();109 }110 111 // Store base pointer extracted from operand into the i-th position of112 // argBase.113 llvm::Value *ptrBaseGEP = builder.CreateInBoundsGEP(114 arrI8PtrTy, mapperAllocas.ArgsBase,115 {builder.getInt32(0), builder.getInt32(index)});116 builder.CreateStore(dataPtrBase, ptrBaseGEP);117 118 // Store pointer extracted from operand into the i-th position of args.119 llvm::Value *ptrGEP = builder.CreateInBoundsGEP(120 arrI8PtrTy, mapperAllocas.Args,121 {builder.getInt32(0), builder.getInt32(index)});122 builder.CreateStore(dataPtr, ptrGEP);123 124 // Store size extracted from operand into the i-th position of argSizes.125 llvm::Value *sizeGEP = builder.CreateInBoundsGEP(126 arrI64Ty, mapperAllocas.ArgSizes,127 {builder.getInt32(0), builder.getInt32(index)});128 builder.CreateStore(dataSize, sizeGEP);129 130 flags.push_back(operandFlag);131 llvm::Constant *mapName =132 mlir::LLVM::createMappingInformation(data.getLoc(), *accBuilder);133 names.push_back(mapName);134 ++index;135 }136 return success();137}138 139/// Process data operands from acc::EnterDataOp140static LogicalResult141processDataOperands(llvm::IRBuilderBase &builder,142 LLVM::ModuleTranslation &moduleTranslation,143 acc::EnterDataOp op, SmallVector<uint64_t> &flags,144 SmallVectorImpl<llvm::Constant *> &names,145 struct OpenACCIRBuilder::MapperAllocas &mapperAllocas) {146 // TODO add `create_zero` and `attach` operands147 148 unsigned index = 0;149 150 // Create operands are handled as `alloc` call.151 // Copyin operands are handled as `to` call.152 llvm::SmallVector<mlir::Value> create, copyin;153 for (mlir::Value dataOp : op.getDataClauseOperands()) {154 if (auto createOp = dataOp.getDefiningOp<acc::CreateOp>()) {155 create.push_back(createOp.getVarPtr());156 } else if (auto copyinOp = mlir::dyn_cast_or_null<acc::CopyinOp>(157 dataOp.getDefiningOp())) {158 copyin.push_back(copyinOp.getVarPtr());159 }160 }161 162 auto nbTotalOperands = create.size() + copyin.size();163 164 // Create operands are handled as `alloc` call.165 if (failed(processOperands(builder, moduleTranslation, op, create,166 nbTotalOperands, kCreateFlag, flags, names, index,167 mapperAllocas)))168 return failure();169 170 // Copyin operands are handled as `to` call.171 if (failed(processOperands(builder, moduleTranslation, op, copyin,172 nbTotalOperands, kDeviceCopyinFlag, flags, names,173 index, mapperAllocas)))174 return failure();175 176 return success();177}178 179/// Process data operands from acc::ExitDataOp180static LogicalResult181processDataOperands(llvm::IRBuilderBase &builder,182 LLVM::ModuleTranslation &moduleTranslation,183 acc::ExitDataOp op, SmallVector<uint64_t> &flags,184 SmallVectorImpl<llvm::Constant *> &names,185 struct OpenACCIRBuilder::MapperAllocas &mapperAllocas) {186 // TODO add `detach` operands187 188 unsigned index = 0;189 190 llvm::SmallVector<mlir::Value> deleteOperands, copyoutOperands;191 for (mlir::Value dataOp : op.getDataClauseOperands()) {192 if (auto devicePtrOp = mlir::dyn_cast_or_null<acc::GetDevicePtrOp>(193 dataOp.getDefiningOp())) {194 for (auto &u : devicePtrOp.getAccPtr().getUses()) {195 if (mlir::dyn_cast_or_null<acc::DeleteOp>(u.getOwner()))196 deleteOperands.push_back(devicePtrOp.getVarPtr());197 else if (mlir::dyn_cast_or_null<acc::CopyoutOp>(u.getOwner()))198 copyoutOperands.push_back(devicePtrOp.getVarPtr());199 }200 }201 }202 203 auto nbTotalOperands = deleteOperands.size() + copyoutOperands.size();204 205 // Delete operands are handled as `delete` call.206 if (failed(processOperands(builder, moduleTranslation, op, deleteOperands,207 nbTotalOperands, kDeleteFlag, flags, names, index,208 mapperAllocas)))209 return failure();210 211 // Copyout operands are handled as `from` call.212 if (failed(processOperands(builder, moduleTranslation, op, copyoutOperands,213 nbTotalOperands, kHostCopyoutFlag, flags, names,214 index, mapperAllocas)))215 return failure();216 217 return success();218}219 220/// Process data operands from acc::UpdateOp221static LogicalResult222processDataOperands(llvm::IRBuilderBase &builder,223 LLVM::ModuleTranslation &moduleTranslation,224 acc::UpdateOp op, SmallVector<uint64_t> &flags,225 SmallVectorImpl<llvm::Constant *> &names,226 struct OpenACCIRBuilder::MapperAllocas &mapperAllocas) {227 unsigned index = 0;228 229 // Host operands are handled as `from` call.230 // Device operands are handled as `to` call.231 llvm::SmallVector<mlir::Value> from, to;232 for (mlir::Value dataOp : op.getDataClauseOperands()) {233 if (auto getDevicePtrOp = mlir::dyn_cast_or_null<acc::GetDevicePtrOp>(234 dataOp.getDefiningOp())) {235 from.push_back(getDevicePtrOp.getVarPtr());236 } else if (auto updateDeviceOp =237 mlir::dyn_cast_or_null<acc::UpdateDeviceOp>(238 dataOp.getDefiningOp())) {239 to.push_back(updateDeviceOp.getVarPtr());240 }241 }242 243 if (failed(processOperands(builder, moduleTranslation, op, from, from.size(),244 kHostCopyoutFlag, flags, names, index,245 mapperAllocas)))246 return failure();247 248 if (failed(processOperands(builder, moduleTranslation, op, to, to.size(),249 kDeviceCopyinFlag, flags, names, index,250 mapperAllocas)))251 return failure();252 return success();253}254 255//===----------------------------------------------------------------------===//256// Conversion functions257//===----------------------------------------------------------------------===//258 259/// Converts an OpenACC data operation into LLVM IR.260static LogicalResult convertDataOp(acc::DataOp &op,261 llvm::IRBuilderBase &builder,262 LLVM::ModuleTranslation &moduleTranslation) {263 llvm::LLVMContext &ctx = builder.getContext();264 auto enclosingFuncOp = op.getOperation()->getParentOfType<LLVM::LLVMFuncOp>();265 llvm::Function *enclosingFunction =266 moduleTranslation.lookupFunction(enclosingFuncOp.getName());267 268 OpenACCIRBuilder *accBuilder = moduleTranslation.getOpenMPBuilder();269 270 llvm::Value *srcLocInfo = createSourceLocationInfo(*accBuilder, op);271 272 llvm::Function *beginMapperFunc = accBuilder->getOrCreateRuntimeFunctionPtr(273 llvm::omp::OMPRTL___tgt_target_data_begin_mapper);274 275 llvm::Function *endMapperFunc = accBuilder->getOrCreateRuntimeFunctionPtr(276 llvm::omp::OMPRTL___tgt_target_data_end_mapper);277 278 // Number of arguments in the data operation.279 unsigned totalNbOperand = op.getNumDataOperands();280 281 struct OpenACCIRBuilder::MapperAllocas mapperAllocas;282 OpenACCIRBuilder::InsertPointTy allocaIP(283 &enclosingFunction->getEntryBlock(),284 enclosingFunction->getEntryBlock().getFirstInsertionPt());285 accBuilder->createMapperAllocas(builder.saveIP(), allocaIP, totalNbOperand,286 mapperAllocas);287 288 SmallVector<uint64_t> flags;289 SmallVector<llvm::Constant *> names;290 unsigned index = 0;291 292 // TODO handle no_create, deviceptr and attach operands.293 294 llvm::SmallVector<mlir::Value> copyin, copyout, create, present,295 deleteOperands;296 for (mlir::Value dataOp : op.getDataClauseOperands()) {297 if (auto devicePtrOp = mlir::dyn_cast_or_null<acc::GetDevicePtrOp>(298 dataOp.getDefiningOp())) {299 for (auto &u : devicePtrOp.getAccPtr().getUses()) {300 if (mlir::dyn_cast_or_null<acc::DeleteOp>(u.getOwner())) {301 deleteOperands.push_back(devicePtrOp.getVarPtr());302 } else if (mlir::dyn_cast_or_null<acc::CopyoutOp>(u.getOwner())) {303 // TODO copyout zero currenlty handled as copyout. Update when304 // extension available.305 copyout.push_back(devicePtrOp.getVarPtr());306 }307 }308 } else if (auto copyinOp = mlir::dyn_cast_or_null<acc::CopyinOp>(309 dataOp.getDefiningOp())) {310 // TODO copyin readonly currenlty handled as copyin. Update when extension311 // available.312 copyin.push_back(copyinOp.getVarPtr());313 } else if (auto createOp = mlir::dyn_cast_or_null<acc::CreateOp>(314 dataOp.getDefiningOp())) {315 // TODO create zero currenlty handled as create. Update when extension316 // available.317 create.push_back(createOp.getVarPtr());318 } else if (auto presentOp = mlir::dyn_cast_or_null<acc::PresentOp>(319 dataOp.getDefiningOp())) {320 present.push_back(createOp.getVarPtr());321 }322 }323 324 auto nbTotalOperands = copyin.size() + copyout.size() + create.size() +325 present.size() + deleteOperands.size();326 327 // Copyin operands are handled as `to` call.328 if (failed(processOperands(builder, moduleTranslation, op, copyin,329 nbTotalOperands, kDeviceCopyinFlag | kHoldFlag,330 flags, names, index, mapperAllocas)))331 return failure();332 333 // Delete operands are handled as `delete` call.334 if (failed(processOperands(builder, moduleTranslation, op, deleteOperands,335 nbTotalOperands, kDeleteFlag, flags, names, index,336 mapperAllocas)))337 return failure();338 339 // Copyout operands are handled as `from` call.340 if (failed(processOperands(builder, moduleTranslation, op, copyout,341 nbTotalOperands, kHostCopyoutFlag | kHoldFlag,342 flags, names, index, mapperAllocas)))343 return failure();344 345 // Create operands are handled as `alloc` call.346 if (failed(processOperands(builder, moduleTranslation, op, create,347 nbTotalOperands, kCreateFlag | kHoldFlag, flags,348 names, index, mapperAllocas)))349 return failure();350 351 if (failed(processOperands(builder, moduleTranslation, op, present,352 nbTotalOperands, kPresentFlag | kHoldFlag, flags,353 names, index, mapperAllocas)))354 return failure();355 356 llvm::GlobalVariable *maptypes =357 accBuilder->createOffloadMaptypes(flags, ".offload_maptypes");358 llvm::Value *maptypesArg = builder.CreateConstInBoundsGEP2_32(359 llvm::ArrayType::get(llvm::Type::getInt64Ty(ctx), totalNbOperand),360 maptypes, /*Idx0=*/0, /*Idx1=*/0);361 362 llvm::GlobalVariable *mapnames =363 accBuilder->createOffloadMapnames(names, ".offload_mapnames");364 llvm::Value *mapnamesArg = builder.CreateConstInBoundsGEP2_32(365 llvm::ArrayType::get(llvm::PointerType::getUnqual(ctx), totalNbOperand),366 mapnames, /*Idx0=*/0, /*Idx1=*/0);367 368 // Create call to start the data region.369 accBuilder->emitMapperCall(builder.saveIP(), beginMapperFunc, srcLocInfo,370 maptypesArg, mapnamesArg, mapperAllocas,371 kDefaultDevice, totalNbOperand);372 373 // Convert the region.374 llvm::BasicBlock *entryBlock = nullptr;375 376 for (Block &bb : op.getRegion()) {377 llvm::BasicBlock *llvmBB = llvm::BasicBlock::Create(378 ctx, "acc.data", builder.GetInsertBlock()->getParent());379 if (entryBlock == nullptr)380 entryBlock = llvmBB;381 moduleTranslation.mapBlock(&bb, llvmBB);382 }383 384 auto afterDataRegion = builder.saveIP();385 386 llvm::BranchInst *sourceTerminator = builder.CreateBr(entryBlock);387 388 builder.restoreIP(afterDataRegion);389 llvm::BasicBlock *endDataBlock = llvm::BasicBlock::Create(390 ctx, "acc.end_data", builder.GetInsertBlock()->getParent());391 392 SetVector<Block *> blocks = getBlocksSortedByDominance(op.getRegion());393 for (Block *bb : blocks) {394 llvm::BasicBlock *llvmBB = moduleTranslation.lookupBlock(bb);395 if (bb->isEntryBlock()) {396 assert(sourceTerminator->getNumSuccessors() == 1 &&397 "provided entry block has multiple successors");398 sourceTerminator->setSuccessor(0, llvmBB);399 }400 401 if (failed(402 moduleTranslation.convertBlock(*bb, bb->isEntryBlock(), builder))) {403 return failure();404 }405 406 if (isa<acc::TerminatorOp, acc::YieldOp>(bb->getTerminator()))407 builder.CreateBr(endDataBlock);408 }409 410 // Create call to end the data region.411 builder.SetInsertPoint(endDataBlock);412 accBuilder->emitMapperCall(builder.saveIP(), endMapperFunc, srcLocInfo,413 maptypesArg, mapnamesArg, mapperAllocas,414 kDefaultDevice, totalNbOperand);415 416 return success();417}418 419/// Converts an OpenACC standalone data operation into LLVM IR.420template <typename OpTy>421static LogicalResult422convertStandaloneDataOp(OpTy &op, llvm::IRBuilderBase &builder,423 LLVM::ModuleTranslation &moduleTranslation) {424 auto enclosingFuncOp =425 op.getOperation()->template getParentOfType<LLVM::LLVMFuncOp>();426 llvm::Function *enclosingFunction =427 moduleTranslation.lookupFunction(enclosingFuncOp.getName());428 429 OpenACCIRBuilder *accBuilder = moduleTranslation.getOpenMPBuilder();430 431 auto *srcLocInfo = createSourceLocationInfo(*accBuilder, op);432 auto *mapperFunc = getAssociatedFunction(*accBuilder, op);433 434 // Number of arguments in the enter_data operation.435 unsigned totalNbOperand = op.getNumDataOperands();436 437 llvm::LLVMContext &ctx = builder.getContext();438 439 struct OpenACCIRBuilder::MapperAllocas mapperAllocas;440 OpenACCIRBuilder::InsertPointTy allocaIP(441 &enclosingFunction->getEntryBlock(),442 enclosingFunction->getEntryBlock().getFirstInsertionPt());443 accBuilder->createMapperAllocas(builder.saveIP(), allocaIP, totalNbOperand,444 mapperAllocas);445 446 SmallVector<uint64_t> flags;447 SmallVector<llvm::Constant *> names;448 449 if (failed(processDataOperands(builder, moduleTranslation, op, flags, names,450 mapperAllocas)))451 return failure();452 453 llvm::GlobalVariable *maptypes =454 accBuilder->createOffloadMaptypes(flags, ".offload_maptypes");455 llvm::Value *maptypesArg = builder.CreateConstInBoundsGEP2_32(456 llvm::ArrayType::get(llvm::Type::getInt64Ty(ctx), totalNbOperand),457 maptypes, /*Idx0=*/0, /*Idx1=*/0);458 459 llvm::GlobalVariable *mapnames =460 accBuilder->createOffloadMapnames(names, ".offload_mapnames");461 llvm::Value *mapnamesArg = builder.CreateConstInBoundsGEP2_32(462 llvm::ArrayType::get(llvm::PointerType::getUnqual(ctx), totalNbOperand),463 mapnames, /*Idx0=*/0, /*Idx1=*/0);464 465 accBuilder->emitMapperCall(builder.saveIP(), mapperFunc, srcLocInfo,466 maptypesArg, mapnamesArg, mapperAllocas,467 kDefaultDevice, totalNbOperand);468 469 return success();470}471 472namespace {473 474/// Implementation of the dialect interface that converts operations belonging475/// to the OpenACC dialect to LLVM IR.476class OpenACCDialectLLVMIRTranslationInterface477 : public LLVMTranslationDialectInterface {478public:479 using LLVMTranslationDialectInterface::LLVMTranslationDialectInterface;480 481 /// Translates the given operation to LLVM IR using the provided IR builder482 /// and saving the state in `moduleTranslation`.483 LogicalResult484 convertOperation(Operation *op, llvm::IRBuilderBase &builder,485 LLVM::ModuleTranslation &moduleTranslation) const final;486};487 488} // namespace489 490/// Given an OpenACC MLIR operation, create the corresponding LLVM IR491/// (including OpenACC runtime calls).492LogicalResult OpenACCDialectLLVMIRTranslationInterface::convertOperation(493 Operation *op, llvm::IRBuilderBase &builder,494 LLVM::ModuleTranslation &moduleTranslation) const {495 496 return llvm::TypeSwitch<Operation *, LogicalResult>(op)497 .Case([&](acc::DataOp dataOp) {498 return convertDataOp(dataOp, builder, moduleTranslation);499 })500 .Case([&](acc::EnterDataOp enterDataOp) {501 return convertStandaloneDataOp<acc::EnterDataOp>(enterDataOp, builder,502 moduleTranslation);503 })504 .Case([&](acc::ExitDataOp exitDataOp) {505 return convertStandaloneDataOp<acc::ExitDataOp>(exitDataOp, builder,506 moduleTranslation);507 })508 .Case([&](acc::UpdateOp updateOp) {509 return convertStandaloneDataOp<acc::UpdateOp>(updateOp, builder,510 moduleTranslation);511 })512 .Case<acc::TerminatorOp, acc::YieldOp>([](auto op) {513 // `yield` and `terminator` can be just omitted. The block structure was514 // created in the function that handles their parent operation.515 assert(op->getNumOperands() == 0 &&516 "unexpected OpenACC terminator with operands");517 return success();518 })519 .Case<acc::CreateOp, acc::CopyinOp, acc::CopyoutOp, acc::DeleteOp,520 acc::UpdateDeviceOp, acc::GetDevicePtrOp>([](auto op) {521 // NOP522 return success();523 })524 .Default([&](Operation *op) {525 return op->emitError("unsupported OpenACC operation: ")526 << op->getName();527 });528}529 530void mlir::registerOpenACCDialectTranslation(DialectRegistry ®istry) {531 registry.insert<acc::OpenACCDialect>();532 registry.addExtension(+[](MLIRContext *ctx, acc::OpenACCDialect *dialect) {533 dialect->addInterfaces<OpenACCDialectLLVMIRTranslationInterface>();534 });535}536 537void mlir::registerOpenACCDialectTranslation(MLIRContext &context) {538 DialectRegistry registry;539 registerOpenACCDialectTranslation(registry);540 context.appendDialectRegistry(registry);541}542