Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 3 additions & 1 deletion clang/lib/CIR/CodeGen/CIRGenModule.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -908,6 +908,7 @@ cir::GlobalOp
CIRGenModule::getOrCreateCIRGlobal(StringRef mangledName, mlir::Type ty,
LangAS langAS, const VarDecl *d,
ForDefinition_t isForDefinition) {

// Lookup the entry, lazily creating it if necessary.
cir::GlobalOp entry;
if (mlir::Operation *v = getGlobalValue(mangledName)) {
Expand All @@ -918,13 +919,14 @@ CIRGenModule::getOrCreateCIRGlobal(StringRef mangledName, mlir::Type ty,
}

if (entry) {
mlir::ptr::MemorySpaceAttrInterface entryCIRAS = entry.getAddrSpaceAttr();
assert(!cir::MissingFeatures::opGlobalWeakRef());

assert(!cir::MissingFeatures::setDLLStorageClass());
assert(!cir::MissingFeatures::openMP());

if (entry.getSymType() == ty &&
(cir::isMatchingAddressSpace(entry.getAddrSpaceAttr(), langAS)))
cir::isMatchingAddressSpace(entryCIRAS, langAS))
return entry;

// If there are two attempts to define the same mangled name, issue an
Expand Down
32 changes: 31 additions & 1 deletion clang/lib/CIR/CodeGen/TargetInfo.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
#include "CIRGenFunction.h"
#include "CIRGenModule.h"
#include "mlir/Dialect/Ptr/IR/MemorySpaceInterfaces.h"
#include "clang/Basic/AddressSpaces.h"
#include "clang/CIR/Dialect/IR/CIRAttrs.h"
#include "clang/CIR/Dialect/IR/CIRDialect.h"

Expand Down Expand Up @@ -70,6 +71,36 @@ class AMDGPUTargetCIRGenInfo : public TargetCIRGenInfo {
}
}
}

clang::LangAS
getGlobalVarAddressSpace(CIRGenModule &cgm,
const clang::VarDecl *decl) const override {
using clang::LangAS;
assert(!cgm.getLangOpts().OpenCL &&
!(cgm.getLangOpts().CUDA && cgm.getLangOpts().CUDAIsDevice) &&
"Address space agnostic languages only");
LangAS defaultGlobalAS = LangAS::opencl_global;
if (!decl)
return defaultGlobalAS;

LangAS addrSpace = decl->getType().getAddressSpace();
if (addrSpace != LangAS::Default)
return addrSpace;

// Only promote to address space 4 if VarDecl has constant initialization.
if (decl->getType().isConstantStorage(cgm.getASTContext(), false, false) &&
decl->hasConstantInitialization())
return LangAS::opencl_constant;

return defaultGlobalAS;
}

mlir::ptr::MemorySpaceAttrInterface
getCIRAllocaAddressSpace() const override {
return cir::LangAddressSpaceAttr::get(
&getABIInfo().cgt.getMLIRContext(),
cir::LangAddressSpace::OffloadPrivate);
}
};

} // namespace
Expand All @@ -86,7 +117,6 @@ class X8664TargetCIRGenInfo : public TargetCIRGenInfo {
X8664TargetCIRGenInfo(CIRGenTypes &cgt)
: TargetCIRGenInfo(std::make_unique<X8664ABIInfo>(cgt)) {}
};

} // namespace

namespace {
Expand Down
262 changes: 261 additions & 1 deletion clang/lib/CIR/Dialect/Transforms/TargetLowering.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -11,10 +11,15 @@
//===----------------------------------------------------------------------===//

#include "TargetLowering/LowerModule.h"
#include "TargetLowering/TargetLoweringInfo.h"

#include "mlir/IR/PatternMatch.h"
#include "mlir/Support/LLVM.h"
#include "mlir/Transforms/DialectConversion.h"
#include "clang/CIR/Dialect/IR/CIRAttrs.h"
#include "clang/CIR/Dialect/IR/CIRDialect.h"
#include "clang/CIR/Dialect/IR/CIRTypes.h"
#include "clang/CIR/Dialect/Passes.h"
#include "llvm/ADT/TypeSwitch.h"

using namespace mlir;
using namespace cir;
Expand All @@ -32,6 +37,168 @@ struct TargetLoweringPass
void runOnOperation() override;
};

/// A generic target lowering pattern that matches any CIR op whose operand or
/// result types need address space conversion. Clones the op with converted
/// types.
class CIRGenericTargetLoweringPattern : public mlir::ConversionPattern {
public:
CIRGenericTargetLoweringPattern(mlir::MLIRContext *context,
const mlir::TypeConverter &typeConverter)
: mlir::ConversionPattern(typeConverter, MatchAnyOpTypeTag(),
/*benefit=*/1, context) {}

mlir::LogicalResult
matchAndRewrite(mlir::Operation *op, llvm::ArrayRef<mlir::Value> operands,
mlir::ConversionPatternRewriter &rewriter) const override {
// Do not match on operations that have dedicated lowering patterns.
if (llvm::isa<cir::FuncOp, cir::GlobalOp>(op))
return mlir::failure();

const mlir::TypeConverter *typeConverter = getTypeConverter();
assert(typeConverter &&
"CIRGenericTargetLoweringPattern requires a type converter");
bool operandsAndResultsLegal = typeConverter->isLegal(op);
bool regionsLegal =
std::all_of(op->getRegions().begin(), op->getRegions().end(),
[typeConverter](mlir::Region &region) {
return typeConverter->isLegal(&region);
});
if (operandsAndResultsLegal && regionsLegal)
return mlir::failure();

assert(op->getNumRegions() == 0 && "CIRGenericTargetLoweringPattern cannot "
"deal with operations with regions");

mlir::OperationState loweredOpState(op->getLoc(), op->getName());
loweredOpState.addOperands(operands);

// Copy attributes, converting any TypeAttr through the type converter so
// that address-space-bearing types (e.g. AllocaOp's allocaType) stay in
// sync with the converted result types.
for (mlir::NamedAttribute attr : op->getAttrs()) {
if (auto typeAttr = mlir::dyn_cast<mlir::TypeAttr>(attr.getValue())) {
mlir::Type converted = typeConverter->convertType(typeAttr.getValue());
loweredOpState.addAttribute(attr.getName(),
mlir::TypeAttr::get(converted));
} else {
loweredOpState.addAttribute(attr.getName(), attr.getValue());
}
}

loweredOpState.addSuccessors(op->getSuccessors());

llvm::SmallVector<mlir::Type> loweredResultTypes;
loweredResultTypes.reserve(op->getNumResults());
for (mlir::Type result : op->getResultTypes())
loweredResultTypes.push_back(typeConverter->convertType(result));
loweredOpState.addTypes(loweredResultTypes);

for (mlir::Region &region : op->getRegions()) {
mlir::Region *loweredRegion = loweredOpState.addRegion();
rewriter.inlineRegionBefore(region, *loweredRegion, loweredRegion->end());
if (mlir::failed(
rewriter.convertRegionTypes(loweredRegion, *getTypeConverter())))
return mlir::failure();
}

mlir::Operation *loweredOp = rewriter.create(loweredOpState);
rewriter.replaceOp(op, loweredOp);
return mlir::success();
}
};

/// Pattern to lower GlobalOp address space attributes. GlobalOp carries
/// addr_space as a standalone attribute (not inside a type), so the
/// TypeConverter won't reach it automatically.
class CIRGlobalOpTargetLowering
: public mlir::OpConversionPattern<cir::GlobalOp> {
const cir::TargetLoweringInfo &targetInfo;

public:
CIRGlobalOpTargetLowering(mlir::MLIRContext *context,
const mlir::TypeConverter &typeConverter,
const cir::TargetLoweringInfo &targetInfo)
: mlir::OpConversionPattern<cir::GlobalOp>(typeConverter, context,
/*benefit=*/1),
targetInfo(targetInfo) {}

mlir::LogicalResult
matchAndRewrite(cir::GlobalOp op, OpAdaptor adaptor,
mlir::ConversionPatternRewriter &rewriter) const override {
mlir::Type loweredSymTy = getTypeConverter()->convertType(op.getSymType());
if (!loweredSymTy)
return mlir::failure();

// Convert the addr_space attribute.
mlir::ptr::MemorySpaceAttrInterface addrSpace = op.getAddrSpaceAttr();
if (auto langAS =
mlir::dyn_cast_if_present<cir::LangAddressSpaceAttr>(addrSpace)) {
unsigned targetAS =
targetInfo.getTargetAddrSpaceFromCIRAddrSpace(langAS.getValue());
addrSpace =
targetAS == 0
? nullptr
: cir::TargetAddressSpaceAttr::get(op.getContext(), targetAS);
}

// Only rewrite if something actually changed.
if (loweredSymTy == op.getSymType() && addrSpace == op.getAddrSpaceAttr())
return mlir::failure();

auto newOp = mlir::cast<cir::GlobalOp>(rewriter.clone(*op.getOperation()));
newOp.setSymType(loweredSymTy);
newOp.setAddrSpaceAttr(addrSpace);
rewriter.replaceOp(op, newOp);
return mlir::success();
}
};

/// Pattern to lower FuncOp types that contain address spaces.
class CIRFuncOpTargetLowering : public mlir::OpConversionPattern<cir::FuncOp> {
public:
using mlir::OpConversionPattern<cir::FuncOp>::OpConversionPattern;

mlir::LogicalResult
matchAndRewrite(cir::FuncOp op, OpAdaptor adaptor,
mlir::ConversionPatternRewriter &rewriter) const override {
cir::FuncType opFuncType = op.getFunctionType();
mlir::TypeConverter::SignatureConversion signatureConversion(
opFuncType.getNumInputs());

for (const auto &[i, argType] : llvm::enumerate(opFuncType.getInputs())) {
mlir::Type loweredArgType = getTypeConverter()->convertType(argType);
if (!loweredArgType)
return mlir::failure();
signatureConversion.addInputs(i, loweredArgType);
}

mlir::Type loweredReturnType =
getTypeConverter()->convertType(opFuncType.getReturnType());
if (!loweredReturnType)
return mlir::failure();

auto loweredFuncType = cir::FuncType::get(
signatureConversion.getConvertedTypes(), loweredReturnType,
/*isVarArg=*/opFuncType.getVarArg());

// Nothing changed, skip.
if (loweredFuncType == opFuncType)
return mlir::failure();

cir::FuncOp loweredFuncOp = rewriter.cloneWithoutRegions(op);
loweredFuncOp.setFunctionType(loweredFuncType);
rewriter.inlineRegionBefore(op.getBody(), loweredFuncOp.getBody(),
loweredFuncOp.end());
if (mlir::failed(rewriter.convertRegionTypes(&loweredFuncOp.getBody(),
*getTypeConverter(),
&signatureConversion)))
return mlir::failure();

rewriter.eraseOp(op);
return mlir::success();
}
};

} // namespace

static void convertSyncScopeIfPresent(mlir::Operation *op,
Expand All @@ -47,6 +214,80 @@ static void convertSyncScopeIfPresent(mlir::Operation *op,
}
}

/// Prepare the type converter for the target lowering pass.
/// Converts LangAddressSpaceAttr → TargetAddressSpaceAttr inside pointer types.
static void
prepareTargetLoweringTypeConverter(mlir::TypeConverter &converter,
const cir::TargetLoweringInfo &targetInfo) {
converter.addConversion([](mlir::Type type) { return type; });

converter.addConversion([&converter,
&targetInfo](cir::PointerType type) -> mlir::Type {
mlir::Type pointee = converter.convertType(type.getPointee());
if (!pointee)
return {};
auto addrSpace = type.getAddrSpace();
if (auto langAS =
mlir::dyn_cast_if_present<cir::LangAddressSpaceAttr>(addrSpace)) {
unsigned targetAS =
targetInfo.getTargetAddrSpaceFromCIRAddrSpace(langAS.getValue());
addrSpace =
targetAS == 0
? nullptr
: cir::TargetAddressSpaceAttr::get(type.getContext(), targetAS);
}
return cir::PointerType::get(type.getContext(), pointee, addrSpace);
});

converter.addConversion([&converter](cir::ArrayType type) -> mlir::Type {
mlir::Type loweredElementType =
converter.convertType(type.getElementType());
if (!loweredElementType)
return {};
return cir::ArrayType::get(loweredElementType, type.getSize());
});

converter.addConversion([&converter](cir::FuncType type) -> mlir::Type {
llvm::SmallVector<mlir::Type> loweredInputTypes;
loweredInputTypes.reserve(type.getNumInputs());
if (mlir::failed(
converter.convertTypes(type.getInputs(), loweredInputTypes)))
return {};

mlir::Type loweredReturnType = converter.convertType(type.getReturnType());
if (!loweredReturnType)
return {};

return cir::FuncType::get(loweredInputTypes, loweredReturnType,
/*isVarArg=*/type.getVarArg());
});
}

static void
populateTargetLoweringConversionTarget(mlir::ConversionTarget &target,
const mlir::TypeConverter &tc) {
target.addLegalOp<mlir::ModuleOp>();

target.addDynamicallyLegalDialect<cir::CIRDialect>(
[&tc](mlir::Operation *op) {
if (!tc.isLegal(op))
return false;
return std::all_of(
op->getRegions().begin(), op->getRegions().end(),
[&tc](mlir::Region &region) { return tc.isLegal(&region); });
});

target.addDynamicallyLegalOp<cir::FuncOp>(
[&tc](cir::FuncOp op) { return tc.isLegal(op.getFunctionType()); });

target.addDynamicallyLegalOp<cir::GlobalOp>([&tc](cir::GlobalOp op) {
if (!tc.isLegal(op.getSymType()))
return false;
return !mlir::isa_and_present<cir::LangAddressSpaceAttr>(
op.getAddrSpaceAttr());
});
}

void TargetLoweringPass::runOnOperation() {
auto mod = mlir::cast<mlir::ModuleOp>(getOperation());
std::unique_ptr<cir::LowerModule> lowerModule = cir::createLowerModule(mod);
Expand All @@ -57,11 +298,30 @@ void TargetLoweringPass::runOnOperation() {
return;
}

const auto &targetInfo = lowerModule->getTargetLoweringInfo();

mod->walk([&](mlir::Operation *op) {
if (mlir::isa<cir::LoadOp, cir::StoreOp, cir::AtomicXchgOp,
cir::AtomicCmpXchgOp, cir::AtomicFetchOp>(op))
convertSyncScopeIfPresent(op, *lowerModule);
});

// Address space conversion: LangAddressSpaceAttr → TargetAddressSpaceAttr.
mlir::TypeConverter typeConverter;
prepareTargetLoweringTypeConverter(typeConverter, targetInfo);

mlir::RewritePatternSet patterns(mod.getContext());
patterns.add<CIRGlobalOpTargetLowering>(mod.getContext(), typeConverter,
targetInfo);
patterns.add<CIRFuncOpTargetLowering>(typeConverter, mod.getContext());
patterns.add<CIRGenericTargetLoweringPattern>(mod.getContext(),
typeConverter);

mlir::ConversionTarget target(*mod.getContext());
populateTargetLoweringConversionTarget(target, typeConverter);

if (failed(mlir::applyPartialConversion(mod, target, std::move(patterns))))
signalPassFailure();
}

std::unique_ptr<Pass> mlir::createTargetLoweringPass() {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@ add_clang_library(MLIRCIRTargetLowering
LowerModule.cpp
LowerItaniumCXXABI.cpp
TargetLoweringInfo.cpp
Targets/AMDGPU.cpp

DEPENDS
clangBasic
Expand Down
Loading