brintos

brintos / llvm-project-archived public Read only

0
0
Text · 7.5 KiB · 796354e Raw
199 lines · cpp
1//===- TranslateRegistration.cpp - hooks to mlir-translate ----------------===//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 from SPIR-V binary module to MLIR SPIR-V10// ModuleOp.11//12//===----------------------------------------------------------------------===//13 14#include "mlir/Dialect/SPIRV/IR/SPIRVDialect.h"15#include "mlir/Dialect/SPIRV/IR/SPIRVOps.h"16#include "mlir/IR/Builders.h"17#include "mlir/IR/Verifier.h"18#include "mlir/Parser/Parser.h"19#include "mlir/Support/FileUtilities.h"20#include "mlir/Target/SPIRV/Deserialization.h"21#include "mlir/Target/SPIRV/Serialization.h"22#include "mlir/Tools/mlir-translate/Translation.h"23#include "llvm/ADT/StringRef.h"24#include "llvm/Support/FileSystem.h"25#include "llvm/Support/MemoryBuffer.h"26#include "llvm/Support/Path.h"27#include "llvm/Support/SourceMgr.h"28#include "llvm/Support/ToolOutputFile.h"29 30using namespace mlir;31 32//===----------------------------------------------------------------------===//33// Deserialization registration34//===----------------------------------------------------------------------===//35 36// Deserializes the SPIR-V binary module stored in the file named as37// `inputFilename` and returns a module containing the SPIR-V module.38static OwningOpRef<Operation *>39deserializeModule(const llvm::MemoryBuffer *input, MLIRContext *context,40                  const spirv::DeserializationOptions &options) {41  context->loadDialect<spirv::SPIRVDialect>();42 43  // Make sure the input stream can be treated as a stream of SPIR-V words44  auto *start = input->getBufferStart();45  auto size = input->getBufferSize();46  if (size % sizeof(uint32_t) != 0) {47    emitError(UnknownLoc::get(context))48        << "SPIR-V binary module must contain integral number of 32-bit words";49    return {};50  }51 52  auto binary = llvm::ArrayRef(reinterpret_cast<const uint32_t *>(start),53                               size / sizeof(uint32_t));54  return spirv::deserialize(binary, context, options);55}56 57namespace mlir {58void registerFromSPIRVTranslation() {59  static llvm::cl::opt<bool> enableControlFlowStructurization(60      "spirv-structurize-control-flow",61      llvm::cl::desc(62          "Enable control flow structurization into `spirv.mlir.selection` and "63          "`spirv.mlir.loop`. This may need to be disabled to support "64          "deserialization of early exits (see #138688)"),65      llvm::cl::init(true));66 67  TranslateToMLIRRegistration fromBinary(68      "deserialize-spirv", "deserializes the SPIR-V module",69      [](llvm::SourceMgr &sourceMgr, MLIRContext *context) {70        assert(sourceMgr.getNumBuffers() == 1 && "expected one buffer");71        return deserializeModule(72            sourceMgr.getMemoryBuffer(sourceMgr.getMainFileID()), context,73            {enableControlFlowStructurization});74      });75}76} // namespace mlir77 78//===----------------------------------------------------------------------===//79// Serialization registration80//===----------------------------------------------------------------------===//81 82static LogicalResult83serializeModule(spirv::ModuleOp moduleOp, raw_ostream &output,84                const spirv::SerializationOptions &options) {85  SmallVector<uint32_t, 0> binary;86  if (failed(spirv::serialize(moduleOp, binary)))87    return failure();88 89  size_t sizeInBytes = binary.size() * sizeof(uint32_t);90 91  output.write(reinterpret_cast<char *>(binary.data()), sizeInBytes);92 93  if (options.saveModuleForValidation) {94    size_t dirSeparator =95        options.validationFilePrefix.find(llvm::sys::path::get_separator());96    // If file prefix includes directory check if that directory exists.97    if (dirSeparator != std::string::npos) {98      llvm::StringRef parentDir =99          llvm::sys::path::parent_path(options.validationFilePrefix);100      if (!llvm::sys::fs::is_directory(parentDir))101        return moduleOp.emitError(102            "validation prefix directory does not exist\n");103    }104 105    SmallString<128> filename;106    int fd = 0;107 108    std::error_code errorCode = llvm::sys::fs::createUniqueFile(109        options.validationFilePrefix + "%%%%%%.spv", fd, filename);110    if (errorCode)111      return moduleOp.emitError("error creating validation output file: ")112             << errorCode.message() << "\n";113 114    llvm::raw_fd_ostream validationOutput(fd, /*shouldClose=*/true);115    validationOutput.write(reinterpret_cast<char *>(binary.data()),116                           sizeInBytes);117    validationOutput.flush();118  }119 120  return mlir::success();121}122 123namespace mlir {124void registerToSPIRVTranslation() {125  static llvm::cl::opt<std::string> validationFilesPrefix(126      "spirv-save-validation-files-with-prefix",127      llvm::cl::desc(128          "When non-empty string is passed each serialized SPIR-V module is "129          "saved to an additional file that starts with the given prefix. This "130          "is used to generate separate binaries for validation, where "131          "`--split-input-file` normally combines all outputs into one. The "132          "one combined output (`-o`) is still written. Created files need to "133          "be removed manually once processed."),134      llvm::cl::init(""));135 136  TranslateFromMLIRRegistration toBinary(137      "serialize-spirv", "serialize SPIR-V dialect",138      [](spirv::ModuleOp moduleOp, raw_ostream &output) {139        return serializeModule(moduleOp, output,140                               {true, false, !validationFilesPrefix.empty(),141                                validationFilesPrefix});142      },143      [](DialectRegistry &registry) {144        registry.insert<spirv::SPIRVDialect>();145      });146}147} // namespace mlir148 149//===----------------------------------------------------------------------===//150// Round-trip registration151//===----------------------------------------------------------------------===//152 153static LogicalResult roundTripModule(spirv::ModuleOp module, bool emitDebugInfo,154                                     raw_ostream &output) {155  SmallVector<uint32_t, 0> binary;156  MLIRContext *context = module->getContext();157 158  spirv::SerializationOptions options;159  options.emitDebugInfo = emitDebugInfo;160  if (failed(spirv::serialize(module, binary, options)))161    return failure();162 163  MLIRContext deserializationContext(context->getDialectRegistry());164  // TODO: we should only load the required dialects instead of all dialects.165  deserializationContext.loadAllAvailableDialects();166  // Then deserialize to get back a SPIR-V module.167  OwningOpRef<spirv::ModuleOp> spirvModule =168      spirv::deserialize(binary, &deserializationContext);169  if (!spirvModule)170    return failure();171  spirvModule->print(output);172 173  return mlir::success();174}175 176namespace mlir {177void registerTestRoundtripSPIRV() {178  TranslateFromMLIRRegistration roundtrip(179      "test-spirv-roundtrip", "test roundtrip in SPIR-V dialect",180      [](spirv::ModuleOp module, raw_ostream &output) {181        return roundTripModule(module, /*emitDebugInfo=*/false, output);182      },183      [](DialectRegistry &registry) {184        registry.insert<spirv::SPIRVDialect>();185      });186}187 188void registerTestRoundtripDebugSPIRV() {189  TranslateFromMLIRRegistration roundtrip(190      "test-spirv-roundtrip-debug", "test roundtrip debug in SPIR-V",191      [](spirv::ModuleOp module, raw_ostream &output) {192        return roundTripModule(module, /*emitDebugInfo=*/true, output);193      },194      [](DialectRegistry &registry) {195        registry.insert<spirv::SPIRVDialect>();196      });197}198} // namespace mlir199