[Mlir-commits] [mlir] [mlir][SPIRV] Lower ptr dialect to PBA addresses (PR #206159)

Igor Wodiany llvmlistbot at llvm.org
Mon Jun 29 02:05:35 PDT 2026


================
@@ -0,0 +1,297 @@
+//===- PtrToSPIRV.cpp - Ptr to SPIR-V dialect conversion -----------------===//
+//
+// 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
+//
+//===----------------------------------------------------------------------===//
+
+#include "mlir/Conversion/PtrToSPIRV/PtrToSPIRV.h"
+
+#include "mlir/Dialect/Ptr/IR/PtrOps.h"
+#include "mlir/Dialect/Ptr/IR/PtrTypes.h"
+#include "mlir/Dialect/SPIRV/IR/SPIRVAttributes.h"
+#include "mlir/Dialect/SPIRV/IR/SPIRVDialect.h"
+#include "mlir/Dialect/SPIRV/IR/SPIRVOps.h"
+#include "mlir/Dialect/SPIRV/Transforms/SPIRVConversion.h"
+#include "mlir/Transforms/DialectConversion.h"
+#include "llvm/ADT/StringRef.h"
+#include <limits>
+
+namespace mlir {
+#define GEN_PASS_DEF_CONVERTPTRTOSPIRVPASS
+#include "mlir/Conversion/Passes.h.inc"
+} // namespace mlir
+
+using namespace mlir;
+
+namespace {
+
+static FailureOr<Type> getAddressType(spirv::TargetEnvAttr targetAttr,
+                                      MLIRContext *context) {
+  spirv::AddressingModel addressingModel =
+      spirv::getAddressingModel(targetAttr, /*use64bitAddress=*/true);
+  if (addressingModel == spirv::AddressingModel::PhysicalStorageBuffer64)
+    return IntegerType::get(context, 64);
+
+  return failure();
+}
+
+static LogicalResult getMemoryAccessAttrs(std::optional<int64_t> alignment,
+                                          Builder &builder,
+                                          spirv::MemoryAccessAttr &accessAttr,
+                                          IntegerAttr &alignmentAttr) {
+  if (!alignment)
+    return success();
+  if (*alignment > std::numeric_limits<uint32_t>::max())
+    return failure();
+
+  accessAttr = spirv::MemoryAccessAttr::get(builder.getContext(),
+                                            spirv::MemoryAccess::Aligned);
+  alignmentAttr = builder.getI32IntegerAttr(*alignment);
+  return success();
+}
+
+static LogicalResult checkSupportedPtrLoad(ptr::LoadOp op,
+                                           PatternRewriter &rewriter) {
+  if (op.getVolatile_() || op.getNontemporal() || op.getInvariant() ||
+      op.getInvariantGroup())
+    return rewriter.notifyMatchFailure(
+        op, "unsupported ptr.load memory operand for SPIR-V lowering");
+  if (op.getOrdering() != ptr::AtomicOrdering::not_atomic)
+    return rewriter.notifyMatchFailure(
+        op, "unsupported atomic ptr.load for SPIR-V lowering");
+  return success();
+}
+
+static LogicalResult checkSupportedPtrStore(ptr::StoreOp op,
+                                            PatternRewriter &rewriter) {
+  if (op.getVolatile_() || op.getNontemporal() || op.getInvariantGroup())
+    return rewriter.notifyMatchFailure(
+        op, "unsupported ptr.store memory operand for SPIR-V lowering");
+  if (op.getOrdering() != ptr::AtomicOrdering::not_atomic)
+    return rewriter.notifyMatchFailure(
+        op, "unsupported atomic ptr.store for SPIR-V lowering");
+  return success();
+}
+
+static FailureOr<Value>
+castAddressToPointeeType(Operation *op, Value address, Type pointeeType,
+                         spirv::StorageClass storageClass, Location loc,
+                         PatternRewriter &rewriter) {
+  if (!isa<IntegerType>(address.getType())) {
+    (void)rewriter.notifyMatchFailure(op, "expected integer address operand");
+    return failure();
+  }
+  if (storageClass != spirv::StorageClass::PhysicalStorageBuffer) {
+    (void)rewriter.notifyMatchFailure(
+        op, "only PhysicalStorageBuffer pointer materialization is supported");
+    return failure();
+  }
+
+  auto typedPtrType = spirv::PointerType::get(pointeeType, storageClass);
+  return spirv::ConvertUToPtrOp::create(rewriter, loc, typedPtrType, address)
+      .getResult();
+}
+
+struct PtrTypeOffsetOpConverter final
+    : public OpConversionPattern<ptr::TypeOffsetOp> {
+  using Base::Base;
+
+  LogicalResult
+  matchAndRewrite(ptr::TypeOffsetOp op, OpAdaptor adaptor,
+                  ConversionPatternRewriter &rewriter) const override {
+    Type convertedType = getTypeConverter()->convertType(op.getType());
+    if (!convertedType)
+      return rewriter.notifyMatchFailure(op, "result type is not convertible");
+
+    llvm::TypeSize typeSize = op.getTypeSize();
+    if (typeSize.isScalable())
+      return rewriter.notifyMatchFailure(op, "scalable type size");
+
+    auto intType = dyn_cast<IntegerType>(convertedType);
+    if (!intType)
+      return rewriter.notifyMatchFailure(
+          op, "converted result type is not an integer");
+
+    auto attr = rewriter.getIntegerAttr(intType, typeSize.getFixedValue());
+    rewriter.replaceOpWithNewOp<spirv::ConstantOp>(op, intType, attr);
+    return success();
+  }
+};
+
+struct PtrAddOpConverter final : public OpConversionPattern<ptr::PtrAddOp> {
+  using Base::Base;
+
+  LogicalResult
+  matchAndRewrite(ptr::PtrAddOp op, OpAdaptor adaptor,
+                  ConversionPatternRewriter &rewriter) const override {
+    Type convertedType = getTypeConverter()->convertType(op.getType());
+    if (!convertedType)
+      return rewriter.notifyMatchFailure(op, "result type is not convertible");
+
+    auto addressType = dyn_cast<IntegerType>(convertedType);
+    if (!addressType)
+      return rewriter.notifyMatchFailure(
+          op, "converted result type is not an integer address");
+    if (adaptor.getBase().getType() != addressType)
+      return rewriter.notifyMatchFailure(
+          op, "converted base pointer type does not match result type");
+
+    Location loc = op.getLoc();
+    Value offset = adaptor.getOffset();
+    if (!isa<IntegerType>(offset.getType()))
+      return rewriter.notifyMatchFailure(op, "offset is not an integer");
+    if (offset.getType() != addressType)
+      offset = spirv::UConvertOp::create(rewriter, loc, addressType, offset);
+
+    rewriter.replaceOpWithNewOp<spirv::IAddOp>(op, addressType,
+                                               adaptor.getBase(), offset);
+    return success();
+  }
+};
+
+struct PtrLoadOpConverter final : public OpConversionPattern<ptr::LoadOp> {
+  PtrLoadOpConverter(const SPIRVTypeConverter &typeConverter,
+                     MLIRContext *context, spirv::StorageClass storageClass)
+      : OpConversionPattern<ptr::LoadOp>(typeConverter, context),
+        storageClass(storageClass) {}
+
+  LogicalResult
+  matchAndRewrite(ptr::LoadOp op, OpAdaptor adaptor,
+                  ConversionPatternRewriter &rewriter) const override {
+    if (failed(checkSupportedPtrLoad(op, rewriter)))
+      return failure();
+
+    Type convertedType = getTypeConverter()->convertType(op.getType());
+    if (!convertedType)
+      return rewriter.notifyMatchFailure(op, "result type is not convertible");
+
+    FailureOr<Value> ptr =
+        castAddressToPointeeType(op, adaptor.getPtr(), convertedType,
+                                 storageClass, op.getLoc(), rewriter);
+    if (failed(ptr))
+      return failure();
+
+    spirv::MemoryAccessAttr accessAttr;
+    IntegerAttr alignmentAttr;
+    if (failed(getMemoryAccessAttrs(op.getAlignment(), rewriter, accessAttr,
+                                    alignmentAttr)))
+      return rewriter.notifyMatchFailure(op, "invalid alignment requirement");
+
+    rewriter.replaceOpWithNewOp<spirv::LoadOp>(op, convertedType, *ptr,
+                                               accessAttr, alignmentAttr);
+    return success();
+  }
+
+private:
+  spirv::StorageClass storageClass;
+};
+
+struct PtrStoreOpConverter final : public OpConversionPattern<ptr::StoreOp> {
+  PtrStoreOpConverter(const SPIRVTypeConverter &typeConverter,
+                      MLIRContext *context, spirv::StorageClass storageClass)
+      : OpConversionPattern<ptr::StoreOp>(typeConverter, context),
+        storageClass(storageClass) {}
+
+  LogicalResult
+  matchAndRewrite(ptr::StoreOp op, OpAdaptor adaptor,
+                  ConversionPatternRewriter &rewriter) const override {
+    if (failed(checkSupportedPtrStore(op, rewriter)))
+      return failure();
+
+    Type valueType = adaptor.getValue().getType();
+    FailureOr<Value> ptr = castAddressToPointeeType(
+        op, adaptor.getPtr(), valueType, storageClass, op.getLoc(), rewriter);
+    if (failed(ptr))
+      return failure();
----------------
IgWod wrote:

Same as above.

https://github.com/llvm/llvm-project/pull/206159


More information about the Mlir-commits mailing list