brintos

brintos / llvm-project-archived public Read only

0
0
Text · 1.5 KiB · 1891741 Raw
42 lines · cpp
1//===--- DialectNVGPU.cpp - Pybind module for NVGPU dialect API support ---===//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-c/Dialect/NVGPU.h"10#include "mlir-c/IR.h"11#include "mlir/Bindings/Python/Nanobind.h"12#include "mlir/Bindings/Python/NanobindAdaptors.h"13 14namespace nb = nanobind;15using namespace llvm;16using namespace mlir;17using namespace mlir::python;18using namespace mlir::python::nanobind_adaptors;19 20static void populateDialectNVGPUSubmodule(const nb::module_ &m) {21  auto nvgpuTensorMapDescriptorType = mlir_type_subclass(22      m, "TensorMapDescriptorType", mlirTypeIsANVGPUTensorMapDescriptorType);23 24  nvgpuTensorMapDescriptorType.def_classmethod(25      "get",26      [](const nb::object &cls, MlirType tensorMemrefType, int swizzle,27         int l2promo, int oobFill, int interleave, MlirContext ctx) {28        return cls(mlirNVGPUTensorMapDescriptorTypeGet(29            ctx, tensorMemrefType, swizzle, l2promo, oobFill, interleave));30      },31      "Gets an instance of TensorMapDescriptorType in the same context",32      nb::arg("cls"), nb::arg("tensor_type"), nb::arg("swizzle"),33      nb::arg("l2promo"), nb::arg("oob_fill"), nb::arg("interleave"),34      nb::arg("ctx") = nb::none());35}36 37NB_MODULE(_mlirDialectsNVGPU, m) {38  m.doc() = "MLIR NVGPU dialect.";39 40  populateDialectNVGPUSubmodule(m);41}42