brintos

brintos / llvm-project-archived public Read only

0
0
Text · 4.0 KiB · ef35c39 Raw
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