[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