112 lines · cpp
1//===- X86VectorDialect.cpp - MLIR X86Vector ops implementation -----------===//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 the X86Vector dialect and its operations.10//11//===----------------------------------------------------------------------===//12 13#include "mlir/Dialect/X86Vector/X86VectorDialect.h"14#include "mlir/Dialect/LLVMIR/LLVMDialect.h"15#include "mlir/IR/Builders.h"16#include "mlir/IR/TypeUtilities.h"17 18using namespace mlir;19 20#include "mlir/Dialect/X86Vector/X86VectorInterfaces.cpp.inc"21 22#include "mlir/Dialect/X86Vector/X86VectorDialect.cpp.inc"23 24void x86vector::X86VectorDialect::initialize() {25 addOperations<26#define GET_OP_LIST27#include "mlir/Dialect/X86Vector/X86Vector.cpp.inc"28 >();29}30 31static Value getMemrefBuffPtr(Location loc, MemRefType type, Value buffer,32 const LLVMTypeConverter &typeConverter,33 RewriterBase &rewriter) {34 MemRefDescriptor memRefDescriptor(buffer);35 return memRefDescriptor.bufferPtr(rewriter, loc, typeConverter, type);36}37 38LogicalResult x86vector::MaskCompressOp::verify() {39 if (getSrc() && getConstantSrc())40 return emitError("cannot use both src and constant_src");41 42 if (getSrc() && (getSrc().getType() != getDst().getType()))43 return emitError("failed to verify that src and dst have same type");44 45 if (getConstantSrc() && (getConstantSrc()->getType() != getDst().getType()))46 return emitError(47 "failed to verify that constant_src and dst have same type");48 49 return success();50}51 52SmallVector<Value> x86vector::MaskCompressOp::getIntrinsicOperands(53 ArrayRef<Value> operands, const LLVMTypeConverter &typeConverter,54 RewriterBase &rewriter) {55 auto loc = getLoc();56 Adaptor adaptor(operands, *this);57 58 auto opType = adaptor.getA().getType();59 Value src;60 if (adaptor.getSrc()) {61 src = adaptor.getSrc();62 } else if (adaptor.getConstantSrc()) {63 src = LLVM::ConstantOp::create(rewriter, loc, opType,64 adaptor.getConstantSrcAttr());65 } else {66 auto zeroAttr = rewriter.getZeroAttr(opType);67 src = LLVM::ConstantOp::create(rewriter, loc, opType, zeroAttr);68 }69 70 return SmallVector<Value>{adaptor.getA(), src, adaptor.getK()};71}72 73SmallVector<Value>74x86vector::DotOp::getIntrinsicOperands(ArrayRef<Value> operands,75 const LLVMTypeConverter &typeConverter,76 RewriterBase &rewriter) {77 SmallVector<Value> intrinsicOperands(operands);78 // Dot product of all elements, broadcasted to all elements.79 Value scale =80 LLVM::ConstantOp::create(rewriter, getLoc(), rewriter.getI8Type(), 0xff);81 intrinsicOperands.push_back(scale);82 83 return intrinsicOperands;84}85 86SmallVector<Value> x86vector::BcstToPackedF32Op::getIntrinsicOperands(87 ArrayRef<Value> operands, const LLVMTypeConverter &typeConverter,88 RewriterBase &rewriter) {89 Adaptor adaptor(operands, *this);90 return {getMemrefBuffPtr(getLoc(), getA().getType(), adaptor.getA(),91 typeConverter, rewriter)};92}93 94SmallVector<Value> x86vector::CvtPackedEvenIndexedToF32Op::getIntrinsicOperands(95 ArrayRef<Value> operands, const LLVMTypeConverter &typeConverter,96 RewriterBase &rewriter) {97 Adaptor adaptor(operands, *this);98 return {getMemrefBuffPtr(getLoc(), getA().getType(), adaptor.getA(),99 typeConverter, rewriter)};100}101 102SmallVector<Value> x86vector::CvtPackedOddIndexedToF32Op::getIntrinsicOperands(103 ArrayRef<Value> operands, const LLVMTypeConverter &typeConverter,104 RewriterBase &rewriter) {105 Adaptor adaptor(operands, *this);106 return {getMemrefBuffPtr(getLoc(), getA().getType(), adaptor.getA(),107 typeConverter, rewriter)};108}109 110#define GET_OP_CLASSES111#include "mlir/Dialect/X86Vector/X86Vector.cpp.inc"112