This change adds a transformation and pass to the NvGPU dialect that attempts to optimize reads/writes from a memref representing GPU shared memory in order to avoid bank conflicts. Given a value representing a shared memory memref, it traverses all reads/writes within the parent op and, subject to suitable conditions, rewrites all last dimension index values such that element locations in the final (col) dimension are given by `newColIdx = col % vecSize + perm[row](col/vecSize,row)` where `perm` is a permutation function indexed by `row` and `vecSize` is the vector access size in elements (currently assumes 128bit vectorized accesses, but this can be made a parameter). This specific transformation can help optimize typical distributed & vectorized accesses common to loading matrix multiplication operands to/from shared memory. Differential Revision: https://reviews.llvm.org/D127457
34 lines
900 B
C++
34 lines
900 B
C++
//===- PassDetail.h - NVGPU Pass class details -----------------*- C++ -*-===//
|
|
//
|
|
// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
|
|
// See https://llvm.org/LICENSE.txt for license information.
|
|
// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
|
|
//
|
|
//===----------------------------------------------------------------------===//
|
|
#ifndef DIALECT_NVGPU_TRANSFORMS_PASSDETAIL_H_
|
|
#define DIALECT_NVGPU_TRANSFORMS_PASSDETAIL_H_
|
|
|
|
#include "mlir/IR/BuiltinOps.h"
|
|
#include "mlir/IR/Dialect.h"
|
|
#include "mlir/Pass/Pass.h"
|
|
|
|
namespace mlir {
|
|
namespace arith {
|
|
class ArithmeticDialect;
|
|
} // namespace arith
|
|
|
|
namespace memref {
|
|
class MemRefDialect;
|
|
} // namespace memref
|
|
|
|
namespace vector {
|
|
class VectorDialect;
|
|
} // namespace vector
|
|
|
|
#define GEN_PASS_CLASSES
|
|
#include "mlir/Dialect/NVGPU/Passes.h.inc"
|
|
|
|
} // namespace mlir
|
|
|
|
#endif // DIALECT_NVGPU_TRANSFORMS_PASSDETAIL_H_
|