brintos

brintos / llvm-project-archived public Read only

0
0
Text · 3.0 KiB · 3831387 Raw
92 lines · cpp
1//===- UBToSPIRV.cpp - UB to SPIRV-V dialect conversion -------------------===//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#include "mlir/Conversion/UBToSPIRV/UBToSPIRV.h"10 11#include "mlir/Dialect/SPIRV/IR/SPIRVDialect.h"12#include "mlir/Dialect/SPIRV/Transforms/SPIRVConversion.h"13#include "mlir/Dialect/UB/IR/UBOps.h"14#include "mlir/Pass/Pass.h"15 16namespace mlir {17#define GEN_PASS_DEF_UBTOSPIRVCONVERSIONPASS18#include "mlir/Conversion/Passes.h.inc"19} // namespace mlir20 21using namespace mlir;22 23namespace {24 25struct PoisonOpLowering final : OpConversionPattern<ub::PoisonOp> {26  using Base::Base;27 28  LogicalResult29  matchAndRewrite(ub::PoisonOp op, OpAdaptor,30                  ConversionPatternRewriter &rewriter) const override {31    Type origType = op.getType();32    Type resType = getTypeConverter()->convertType(origType);33    if (!resType)34      return rewriter.notifyMatchFailure(op, [&](Diagnostic &diag) {35        diag << "failed to convert result type " << origType;36      });37 38    rewriter.replaceOpWithNewOp<spirv::UndefOp>(op, resType);39    return success();40  }41};42 43struct UnreachableOpLowering final : OpConversionPattern<ub::UnreachableOp> {44  using Base::Base;45 46  LogicalResult47  matchAndRewrite(ub::UnreachableOp op, OpAdaptor,48                  ConversionPatternRewriter &rewriter) const override {49    rewriter.replaceOpWithNewOp<spirv::UnreachableOp>(op);50    return success();51  }52};53 54} // namespace55 56//===----------------------------------------------------------------------===//57// Pass Definition58//===----------------------------------------------------------------------===//59 60namespace {61struct UBToSPIRVConversionPass final62    : impl::UBToSPIRVConversionPassBase<UBToSPIRVConversionPass> {63  using Base::Base;64 65  void runOnOperation() override {66    Operation *op = getOperation();67    spirv::TargetEnvAttr targetAttr = spirv::lookupTargetEnvOrDefault(op);68    std::unique_ptr<SPIRVConversionTarget> target =69        SPIRVConversionTarget::get(targetAttr);70 71    SPIRVConversionOptions options;72    SPIRVTypeConverter typeConverter(targetAttr, options);73 74    RewritePatternSet patterns(&getContext());75    ub::populateUBToSPIRVConversionPatterns(typeConverter, patterns);76 77    if (failed(applyPartialConversion(op, *target, std::move(patterns))))78      signalPassFailure();79  }80};81} // namespace82 83//===----------------------------------------------------------------------===//84// Pattern Population85//===----------------------------------------------------------------------===//86 87void mlir::ub::populateUBToSPIRVConversionPatterns(88    const SPIRVTypeConverter &converter, RewritePatternSet &patterns) {89  patterns.add<PoisonOpLowering, UnreachableOpLowering>(converter,90                                                        patterns.getContext());91}92