[MLIR][SPIRVToLLVM] Conversion patterns for GLSL ops

This is the first patch that adds support for GLSL extended
instruction set ops. These are direct conversions, apart from `spv.Tan`
that is lowered to `sin() / cos()`.

Reviewed By: antiagainst

Differential Revision: https://reviews.llvm.org/D84627
This commit is contained in:
George Mitenkov 2020-07-30 09:22:44 +03:00
parent cd4e8d7f6f
commit 647e9a54c7
2 changed files with 133 additions and 0 deletions

View File

@ -802,6 +802,25 @@ public:
}
};
class TanPattern : public SPIRVToLLVMConversion<spirv::GLSLTanOp> {
public:
using SPIRVToLLVMConversion<spirv::GLSLTanOp>::SPIRVToLLVMConversion;
LogicalResult
matchAndRewrite(spirv::GLSLTanOp tanOp, ArrayRef<Value> operands,
ConversionPatternRewriter &rewriter) const override {
auto dstType = typeConverter.convertType(tanOp.getType());
if (!dstType)
return failure();
Location loc = tanOp.getLoc();
Value sin = rewriter.create<LLVM::SinOp>(loc, dstType, tanOp.operand());
Value cos = rewriter.create<LLVM::CosOp>(loc, dstType, tanOp.operand());
rewriter.replaceOpWithNewOp<LLVM::FDivOp>(tanOp, dstType, sin, cos);
return success();
}
};
class VariablePattern : public SPIRVToLLVMConversion<spirv::VariableOp> {
public:
using SPIRVToLLVMConversion<spirv::VariableOp>::SPIRVToLLVMConversion;
@ -1026,6 +1045,15 @@ void mlir::populateSPIRVToLLVMConversionPatterns(
// Function Call op
FunctionCallPattern,
// GLSL extended instruction set ops
DirectConversionPattern<spirv::GLSLCeilOp, LLVM::FCeilOp>,
DirectConversionPattern<spirv::GLSLCosOp, LLVM::CosOp>,
DirectConversionPattern<spirv::GLSLExpOp, LLVM::ExpOp>,
DirectConversionPattern<spirv::GLSLFAbsOp, LLVM::FAbsOp>,
DirectConversionPattern<spirv::GLSLLogOp, LLVM::LogOp>,
DirectConversionPattern<spirv::GLSLSinOp, LLVM::SinOp>,
DirectConversionPattern<spirv::GLSLSqrtOp, LLVM::SqrtOp>, TanPattern,
// Logical ops
DirectConversionPattern<spirv::LogicalAndOp, LLVM::AndOp>,
DirectConversionPattern<spirv::LogicalOrOp, LLVM::OrOp>,

View File

@ -0,0 +1,105 @@
// RUN: mlir-opt -convert-spirv-to-llvm %s | FileCheck %s
//===----------------------------------------------------------------------===//
// spv.GLSL.Ceil
//===----------------------------------------------------------------------===//
// CHECK-LABEL: @ceil
func @ceil(%arg0: f32, %arg1: vector<3xf16>) {
// CHECK: "llvm.intr.ceil"(%{{.*}}) : (!llvm.float) -> !llvm.float
%0 = spv.GLSL.Ceil %arg0 : f32
// CHECK: "llvm.intr.ceil"(%{{.*}}) : (!llvm<"<3 x half>">) -> !llvm<"<3 x half>">
%1 = spv.GLSL.Ceil %arg1 : vector<3xf16>
return
}
//===----------------------------------------------------------------------===//
// spv.GLSL.Cos
//===----------------------------------------------------------------------===//
// CHECK-LABEL: @cos
func @cos(%arg0: f32, %arg1: vector<3xf16>) {
// CHECK: "llvm.intr.cos"(%{{.*}}) : (!llvm.float) -> !llvm.float
%0 = spv.GLSL.Cos %arg0 : f32
// CHECK: "llvm.intr.cos"(%{{.*}}) : (!llvm<"<3 x half>">) -> !llvm<"<3 x half>">
%1 = spv.GLSL.Cos %arg1 : vector<3xf16>
return
}
//===----------------------------------------------------------------------===//
// spv.GLSL.Exp
//===----------------------------------------------------------------------===//
// CHECK-LABEL: @exp
func @exp(%arg0: f32, %arg1: vector<3xf16>) {
// CHECK: "llvm.intr.exp"(%{{.*}}) : (!llvm.float) -> !llvm.float
%0 = spv.GLSL.Exp %arg0 : f32
// CHECK: "llvm.intr.exp"(%{{.*}}) : (!llvm<"<3 x half>">) -> !llvm<"<3 x half>">
%1 = spv.GLSL.Exp %arg1 : vector<3xf16>
return
}
//===----------------------------------------------------------------------===//
// spv.GLSL.FAbs
//===----------------------------------------------------------------------===//
// CHECK-LABEL: @fabs
func @fabs(%arg0: f32, %arg1: vector<3xf16>) {
// CHECK: "llvm.intr.fabs"(%{{.*}}) : (!llvm.float) -> !llvm.float
%0 = spv.GLSL.FAbs %arg0 : f32
// CHECK: "llvm.intr.fabs"(%{{.*}}) : (!llvm<"<3 x half>">) -> !llvm<"<3 x half>">
%1 = spv.GLSL.FAbs %arg1 : vector<3xf16>
return
}
//===----------------------------------------------------------------------===//
// spv.GLSL.Log
//===----------------------------------------------------------------------===//
// CHECK-LABEL: @log
func @log(%arg0: f32, %arg1: vector<3xf16>) {
// CHECK: "llvm.intr.log"(%{{.*}}) : (!llvm.float) -> !llvm.float
%0 = spv.GLSL.Log %arg0 : f32
// CHECK: "llvm.intr.log"(%{{.*}}) : (!llvm<"<3 x half>">) -> !llvm<"<3 x half>">
%1 = spv.GLSL.Log %arg1 : vector<3xf16>
return
}
//===----------------------------------------------------------------------===//
// spv.GLSL.Sin
//===----------------------------------------------------------------------===//
// CHECK-LABEL: @sin
func @sin(%arg0: f32, %arg1: vector<3xf16>) {
// CHECK: "llvm.intr.sin"(%{{.*}}) : (!llvm.float) -> !llvm.float
%0 = spv.GLSL.Sin %arg0 : f32
// CHECK: "llvm.intr.sin"(%{{.*}}) : (!llvm<"<3 x half>">) -> !llvm<"<3 x half>">
%1 = spv.GLSL.Sin %arg1 : vector<3xf16>
return
}
//===----------------------------------------------------------------------===//
// spv.GLSL.Sqrt
//===----------------------------------------------------------------------===//
// CHECK-LABEL: @sqrt
func @sqrt(%arg0: f32, %arg1: vector<3xf16>) {
// CHECK: "llvm.intr.sqrt"(%{{.*}}) : (!llvm.float) -> !llvm.float
%0 = spv.GLSL.Sqrt %arg0 : f32
// CHECK: "llvm.intr.sqrt"(%{{.*}}) : (!llvm<"<3 x half>">) -> !llvm<"<3 x half>">
%1 = spv.GLSL.Sqrt %arg1 : vector<3xf16>
return
}
//===----------------------------------------------------------------------===//
// spv.GLSL.Tan
//===----------------------------------------------------------------------===//
// CHECK-LABEL: @tan
func @tan(%arg0: f32) {
// CHECK: %[[SIN:.*]] = "llvm.intr.sin"(%{{.*}}) : (!llvm.float) -> !llvm.float
// CHECK: %[[COS:.*]] = "llvm.intr.cos"(%{{.*}}) : (!llvm.float) -> !llvm.float
// CHECK: llvm.fdiv %[[SIN]], %[[COS]] : !llvm.float
%0 = spv.GLSL.Tan %arg0 : f32
return
}