[llvm] 239ca8d - [IR] Add elementwise modifier to atomicrmw (#189517)
via llvm-commits
llvm-commits at lists.llvm.org
Mon May 4 20:00:51 PDT 2026
Author: Yonah Goldberg
Date: 2026-05-04T20:00:46-07:00
New Revision: 239ca8d6c13a5c05a2363247e02edf29d2f253a2
URL: https://github.com/llvm/llvm-project/commit/239ca8d6c13a5c05a2363247e02edf29d2f253a2
DIFF: https://github.com/llvm/llvm-project/commit/239ca8d6c13a5c05a2363247e02edf29d2f253a2.diff
LOG: [IR] Add elementwise modifier to atomicrmw (#189517)
This PR implements the IR side modifications of [[RFC] Add elementwise
modifier to atomicrmw](https://discourse.llvm.org/t/rfc-add-elementwise-modifier-to-atomicrmw/90134).
Design Decisions:
- In the IR, the current atomicrmw record layout looks like: [ptrty,
ptr, valty, val, operation, vol, ordering, syncscope, align]. To encode
elementwise, I decided to pack it into the operation field, which also
contains the math op (i.e. fadd, fmin, add etc...). I could have changed
the record structure, but that would be slightly more complicated.
- elementwise vector atomics can be vectors of integers because we can always scalarize legally
- elementwise vector atomics need to have power of 2 size. We can potentially remove this restriction later.
Assisted by AI.
Added:
llvm/test/Assembler/invalid-atomicrmw-elementwise.ll
llvm/test/Bitcode/atomicrmw-elementwise.ll
Modified:
llvm/docs/LangRef.rst
llvm/include/llvm/AsmParser/LLToken.h
llvm/include/llvm/Bitcode/LLVMBitCodes.h
llvm/include/llvm/IR/IRBuilder.h
llvm/include/llvm/IR/Instructions.h
llvm/lib/AsmParser/LLLexer.cpp
llvm/lib/AsmParser/LLParser.cpp
llvm/lib/Bitcode/Reader/BitcodeReader.cpp
llvm/lib/Bitcode/Writer/BitcodeWriter.cpp
llvm/lib/IR/AsmWriter.cpp
llvm/lib/IR/Instruction.cpp
llvm/lib/IR/Instructions.cpp
llvm/lib/IR/Verifier.cpp
llvm/test/Assembler/atomic.ll
llvm/test/Bitcode/compatibility.ll
llvm/unittests/Analysis/AliasAnalysisTest.cpp
llvm/unittests/IR/VerifierTest.cpp
Removed:
################################################################################
diff --git a/llvm/docs/LangRef.rst b/llvm/docs/LangRef.rst
index 912cca72e1ad5..88f5028ea4408 100644
--- a/llvm/docs/LangRef.rst
+++ b/llvm/docs/LangRef.rst
@@ -12259,7 +12259,7 @@ Syntax:
::
- atomicrmw [volatile] <operation> ptr <pointer>, <ty> <value> [syncscope("<target-scope>")] <ordering>[, align <alignment>] ; yields ty
+ atomicrmw [volatile] [elementwise] <operation> ptr <pointer>, <ty> <value> [syncscope("<target-scope>")] <ordering>[, align <alignment>] ; yields ty
Overview:
"""""""""
@@ -12297,14 +12297,12 @@ operation. The operation must be one of the following keywords:
- usub_cond
- usub_sat
-For most of these operations, the type of '<value>' must be an integer
-type whose bit width is a power of two greater than or equal to eight.
-For xchg, this
-may also be a floating point or a pointer type with the same size constraints
-as integers. For fadd/fsub/fmax/fmin/fmaximum/fminimum/fmaximumnum/fminimumnum, this must be a floating-point
-or fixed vector of floating-point type. The type of the '``<pointer>``'
-operand must be a pointer to that type. If the ``atomicrmw`` is marked
-as ``volatile``, then the optimizer is not allowed to modify the
+For all of these operations, the type of '<value>' must be a type whose bit width is a power of two greater than or equal to eight.
+For add/sub/and/nand/or/xor/max/min/umax/umin/uinc_wrap/udec_wrap/usub_cond/usub_sat, this must be an integer type, or, if the ``elementwise`` modifier is present, a fixed vector of integer type.
+For fadd/fsub/fmax/fmin/fmaximum/fminimum/fmaximumnum/fminimumnum, this must be a floating-point or fixed vector of floating-point type.
+For xchg, this must be an integer type, floating-point type, or pointer type, or, if the ``elementwise`` modifier is present, a fixed vector of integer type, floating-point type, or pointer type.
+The type of the '<pointer>' operand must be a pointer to the type of '<value>'.
+If the ``atomicrmw`` is marked as ``volatile``, then the optimizer is not allowed to modify the
number or order of execution of this ``atomicrmw`` with other
:ref:`volatile operations <volatile>`.
@@ -12321,6 +12319,10 @@ isn't specified.
An ``atomicrmw`` instruction can also take an optional
":ref:`syncscope <syncscope>`" argument.
+If the ``elementwise`` modifier is present, the instruction has per-element vector
+atomic semantics. It behaves as if it were expanded into one scalar ``atomicrmw`` per element, that are not ordered with respect to each other.
+Without ``elementwise``, vector ``atomicrmw`` keeps whole-value atomic semantics.
+
Semantics:
""""""""""
diff --git a/llvm/include/llvm/AsmParser/LLToken.h b/llvm/include/llvm/AsmParser/LLToken.h
index 593f66900e613..a9864cfbcab25 100644
--- a/llvm/include/llvm/AsmParser/LLToken.h
+++ b/llvm/include/llvm/AsmParser/LLToken.h
@@ -91,6 +91,7 @@ enum Kind {
kw_unwind,
kw_datalayout,
kw_volatile,
+ kw_elementwise,
kw_atomic,
kw_unordered,
kw_monotonic,
diff --git a/llvm/include/llvm/Bitcode/LLVMBitCodes.h b/llvm/include/llvm/Bitcode/LLVMBitCodes.h
index 9162754bbfe1a..95787c595dff7 100644
--- a/llvm/include/llvm/Bitcode/LLVMBitCodes.h
+++ b/llvm/include/llvm/Bitcode/LLVMBitCodes.h
@@ -524,6 +524,10 @@ enum RMWOperations {
RMW_FMINIMUMNUM = 22,
};
+enum RMWOperationFlags {
+ RMW_ELEMENTWISE_FLAG = 1 << 5,
+};
+
/// OverflowingBinaryOperatorOptionalFlags - Flags for serializing
/// OverflowingBinaryOperator's SubclassOptionalData contents.
enum OverflowingBinaryOperatorOptionalFlags {
diff --git a/llvm/include/llvm/IR/IRBuilder.h b/llvm/include/llvm/IR/IRBuilder.h
index 84a588023826d..d3c782a7b4a48 100644
--- a/llvm/include/llvm/IR/IRBuilder.h
+++ b/llvm/include/llvm/IR/IRBuilder.h
@@ -1977,13 +1977,15 @@ class IRBuilderBase {
AtomicRMWInst *CreateAtomicRMW(AtomicRMWInst::BinOp Op, Value *Ptr,
Value *Val, MaybeAlign Align,
AtomicOrdering Ordering,
- SyncScope::ID SSID = SyncScope::System) {
+ SyncScope::ID SSID = SyncScope::System,
+ bool Elementwise = false) {
if (!Align) {
const DataLayout &DL = BB->getDataLayout();
Align = llvm::Align(DL.getTypeStoreSize(Val->getType()));
}
- return Insert(new AtomicRMWInst(Op, Ptr, Val, *Align, Ordering, SSID));
+ return Insert(
+ new AtomicRMWInst(Op, Ptr, Val, *Align, Ordering, SSID, Elementwise));
}
CallInst *CreateStructuredGEP(Type *BaseType, Value *PtrBase,
diff --git a/llvm/include/llvm/IR/Instructions.h b/llvm/include/llvm/IR/Instructions.h
index f4fa59ada9d05..0a6408d832c5b 100644
--- a/llvm/include/llvm/IR/Instructions.h
+++ b/llvm/include/llvm/IR/Instructions.h
@@ -809,7 +809,7 @@ class AtomicRMWInst : public Instruction {
public:
LLVM_ABI AtomicRMWInst(BinOp Operation, Value *Ptr, Value *Val,
Align Alignment, AtomicOrdering Ordering,
- SyncScope::ID SSID,
+ SyncScope::ID SSID, bool Elementwise = false,
InsertPosition InsertBefore = nullptr);
// allocate space for exactly two operands
@@ -821,8 +821,10 @@ class AtomicRMWInst : public Instruction {
AtomicOrderingBitfieldElementT<VolatileField::NextBit>;
using OperationField = BinOpBitfieldElement<AtomicOrderingField::NextBit>;
using AlignmentField = AlignmentBitfieldElementT<OperationField::NextBit>;
+ using ElementwiseField = BoolBitfieldElementT<AlignmentField::NextBit>;
static_assert(Bitfield::areContiguous<VolatileField, AtomicOrderingField,
- OperationField, AlignmentField>(),
+ OperationField, AlignmentField,
+ ElementwiseField>(),
"Bitfields must be contiguous");
BinOp getOperation() const { return getSubclassData<OperationField>(); }
@@ -867,6 +869,12 @@ class AtomicRMWInst : public Instruction {
///
void setVolatile(bool V) { setSubclassData<VolatileField>(V); }
+ /// Return true if this RMW has elementwise vector semantics.
+ bool isElementwise() const { return getSubclassData<ElementwiseField>(); }
+
+ /// Specify whether this RMW has elementwise vector semantics.
+ void setElementwise(bool V) { setSubclassData<ElementwiseField>(V); }
+
/// Transparently provide more efficient getOperand methods.
DECLARE_TRANSPARENT_OPERAND_ACCESSORS(Value);
@@ -920,7 +928,7 @@ class AtomicRMWInst : public Instruction {
private:
void Init(BinOp Operation, Value *Ptr, Value *Val, Align Align,
- AtomicOrdering Ordering, SyncScope::ID SSID);
+ AtomicOrdering Ordering, SyncScope::ID SSID, bool Elementwise);
// Shadow Instruction::setInstructionSubclassData with a private forwarding
// method so that subclasses cannot accidentally use it.
diff --git a/llvm/lib/AsmParser/LLLexer.cpp b/llvm/lib/AsmParser/LLLexer.cpp
index 8084e90271955..e7d7cc4ba88ea 100644
--- a/llvm/lib/AsmParser/LLLexer.cpp
+++ b/llvm/lib/AsmParser/LLLexer.cpp
@@ -601,6 +601,7 @@ lltok::Kind LLLexer::LexIdentifier() {
KEYWORD(unwind);
KEYWORD(datalayout);
KEYWORD(volatile);
+ KEYWORD(elementwise);
KEYWORD(atomic);
KEYWORD(unordered);
KEYWORD(monotonic);
diff --git a/llvm/lib/AsmParser/LLParser.cpp b/llvm/lib/AsmParser/LLParser.cpp
index dadbc510cdbd5..820f64bf30ba7 100644
--- a/llvm/lib/AsmParser/LLParser.cpp
+++ b/llvm/lib/AsmParser/LLParser.cpp
@@ -9002,20 +9002,24 @@ int LLParser::parseCmpXchg(Instruction *&Inst, PerFunctionState &PFS) {
}
/// parseAtomicRMW
-/// ::= 'atomicrmw' 'volatile'? BinOp TypeAndValue ',' TypeAndValue
+/// ::= 'atomicrmw' 'volatile'? 'elementwise'? BinOp TypeAndValue ','
+/// TypeAndValue
/// 'singlethread'? AtomicOrdering
int LLParser::parseAtomicRMW(Instruction *&Inst, PerFunctionState &PFS) {
Value *Ptr, *Val; LocTy PtrLoc, ValLoc;
bool AteExtraComma = false;
AtomicOrdering Ordering = AtomicOrdering::NotAtomic;
SyncScope::ID SSID = SyncScope::System;
- bool isVolatile = false;
+ bool IsVolatile = false;
+ bool IsElementwise = false;
bool IsFP = false;
AtomicRMWInst::BinOp Operation;
MaybeAlign Alignment;
if (EatIfPresent(lltok::kw_volatile))
- isVolatile = true;
+ IsVolatile = true;
+ if (EatIfPresent(lltok::kw_elementwise))
+ IsElementwise = true;
switch (Lex.getKind()) {
default:
@@ -9092,23 +9096,34 @@ int LLParser::parseAtomicRMW(Instruction *&Inst, PerFunctionState &PFS) {
if (Val->getType()->isScalableTy())
return error(ValLoc, "atomicrmw operand may not be scalable");
+ // For elementwise ops, the value must be a fixed vector type whose element
+ // type is legal for the corresponding scalar atomicrmw operation. So assign
+ // ScalarTy the element type for elementwise ops so we can check this.
+ Type *ScalarTy = Val->getType();
+ if (IsElementwise) {
+ auto *VecTy = dyn_cast<FixedVectorType>(Val->getType());
+ if (!VecTy)
+ return error(ValLoc,
+ "atomicrmw elementwise operand must be a fixed vector type");
+ ScalarTy = VecTy->getElementType();
+ }
+
if (Operation == AtomicRMWInst::Xchg) {
- if (!Val->getType()->isIntegerTy() &&
- !Val->getType()->isFloatingPointTy() &&
- !Val->getType()->isPointerTy()) {
+ if (!ScalarTy->isIntegerTy() && !ScalarTy->isFloatingPointTy() &&
+ !ScalarTy->isPointerTy()) {
return error(
ValLoc,
"atomicrmw " + AtomicRMWInst::getOperationName(Operation) +
" operand must be an integer, floating point, or pointer type");
}
} else if (IsFP) {
- if (!Val->getType()->isFPOrFPVectorTy()) {
+ if (!ScalarTy->isFPOrFPVectorTy()) {
return error(ValLoc, "atomicrmw " +
AtomicRMWInst::getOperationName(Operation) +
" operand must be a floating point type");
}
} else {
- if (!Val->getType()->isIntegerTy()) {
+ if (!ScalarTy->isIntegerTy()) {
return error(ValLoc, "atomicrmw " +
AtomicRMWInst::getOperationName(Operation) +
" operand must be an integer");
@@ -9116,18 +9131,16 @@ int LLParser::parseAtomicRMW(Instruction *&Inst, PerFunctionState &PFS) {
}
unsigned Size =
- PFS.getFunction().getDataLayout().getTypeStoreSizeInBits(
- Val->getType());
+ PFS.getFunction().getDataLayout().getTypeStoreSizeInBits(Val->getType());
if (Size < 8 || (Size & (Size - 1)))
- return error(ValLoc, "atomicrmw operand must be power-of-two byte-sized"
- " integer");
+ return error(ValLoc,
+ "atomicrmw operand must have a power-of-two byte size");
const Align DefaultAlignment(
- PFS.getFunction().getDataLayout().getTypeStoreSize(
- Val->getType()));
- AtomicRMWInst *RMWI =
- new AtomicRMWInst(Operation, Ptr, Val,
- Alignment.value_or(DefaultAlignment), Ordering, SSID);
- RMWI->setVolatile(isVolatile);
+ PFS.getFunction().getDataLayout().getTypeStoreSize(Val->getType()));
+ AtomicRMWInst *RMWI = new AtomicRMWInst(Operation, Ptr, Val,
+ Alignment.value_or(DefaultAlignment),
+ Ordering, SSID, IsElementwise);
+ RMWI->setVolatile(IsVolatile);
Inst = RMWI;
return AteExtraComma ? InstExtraComma : InstNormal;
}
diff --git a/llvm/lib/Bitcode/Reader/BitcodeReader.cpp b/llvm/lib/Bitcode/Reader/BitcodeReader.cpp
index fa7a3b214e463..83babe1c62541 100644
--- a/llvm/lib/Bitcode/Reader/BitcodeReader.cpp
+++ b/llvm/lib/Bitcode/Reader/BitcodeReader.cpp
@@ -1365,8 +1365,10 @@ static int getDecodedBinaryOpcode(unsigned Val, Type *Ty) {
}
}
-static AtomicRMWInst::BinOp getDecodedRMWOperation(unsigned Val) {
- switch (Val) {
+static AtomicRMWInst::BinOp getDecodedRMWOperation(unsigned Val,
+ bool &IsElementwise) {
+ IsElementwise = Val & bitc::RMW_ELEMENTWISE_FLAG;
+ switch (Val & ~bitc::RMW_ELEMENTWISE_FLAG) {
default: return AtomicRMWInst::BAD_BINOP;
case bitc::RMW_XCHG: return AtomicRMWInst::Xchg;
case bitc::RMW_ADD: return AtomicRMWInst::Add;
@@ -6647,8 +6649,9 @@ Error BitcodeReader::parseFunctionBody(Function *F) {
if (!(NumRecords == (OpNum + 4) || NumRecords == (OpNum + 5)))
return error("Invalid atomicrmw record");
+ bool IsElementwise = false;
const AtomicRMWInst::BinOp Operation =
- getDecodedRMWOperation(Record[OpNum]);
+ getDecodedRMWOperation(Record[OpNum], IsElementwise);
if (Operation < AtomicRMWInst::FIRST_BINOP ||
Operation > AtomicRMWInst::LAST_BINOP)
return error("Invalid atomicrmw record");
@@ -6673,7 +6676,8 @@ Error BitcodeReader::parseFunctionBody(Function *F) {
Alignment =
Align(TheModule->getDataLayout().getTypeStoreSize(Val->getType()));
- I = new AtomicRMWInst(Operation, Ptr, Val, *Alignment, Ordering, SSID);
+ I = new AtomicRMWInst(Operation, Ptr, Val, *Alignment, Ordering, SSID,
+ IsElementwise);
ResTypeID = ValTypeID;
cast<AtomicRMWInst>(I)->setVolatile(IsVol);
diff --git a/llvm/lib/Bitcode/Writer/BitcodeWriter.cpp b/llvm/lib/Bitcode/Writer/BitcodeWriter.cpp
index 31b8367083087..ed7f95701ea65 100644
--- a/llvm/lib/Bitcode/Writer/BitcodeWriter.cpp
+++ b/llvm/lib/Bitcode/Writer/BitcodeWriter.cpp
@@ -690,41 +690,84 @@ static unsigned getEncodedBinaryOpcode(unsigned Opcode) {
}
}
-static unsigned getEncodedRMWOperation(AtomicRMWInst::BinOp Op) {
- switch (Op) {
+static unsigned getEncodedRMWOperation(const AtomicRMWInst &I) {
+ unsigned Encoding = 0;
+ switch (I.getOperation()) {
default: llvm_unreachable("Unknown RMW operation!");
- case AtomicRMWInst::Xchg: return bitc::RMW_XCHG;
- case AtomicRMWInst::Add: return bitc::RMW_ADD;
- case AtomicRMWInst::Sub: return bitc::RMW_SUB;
- case AtomicRMWInst::And: return bitc::RMW_AND;
- case AtomicRMWInst::Nand: return bitc::RMW_NAND;
- case AtomicRMWInst::Or: return bitc::RMW_OR;
- case AtomicRMWInst::Xor: return bitc::RMW_XOR;
- case AtomicRMWInst::Max: return bitc::RMW_MAX;
- case AtomicRMWInst::Min: return bitc::RMW_MIN;
- case AtomicRMWInst::UMax: return bitc::RMW_UMAX;
- case AtomicRMWInst::UMin: return bitc::RMW_UMIN;
- case AtomicRMWInst::FAdd: return bitc::RMW_FADD;
- case AtomicRMWInst::FSub: return bitc::RMW_FSUB;
- case AtomicRMWInst::FMax: return bitc::RMW_FMAX;
- case AtomicRMWInst::FMin: return bitc::RMW_FMIN;
+ case AtomicRMWInst::Xchg:
+ Encoding = bitc::RMW_XCHG;
+ break;
+ case AtomicRMWInst::Add:
+ Encoding = bitc::RMW_ADD;
+ break;
+ case AtomicRMWInst::Sub:
+ Encoding = bitc::RMW_SUB;
+ break;
+ case AtomicRMWInst::And:
+ Encoding = bitc::RMW_AND;
+ break;
+ case AtomicRMWInst::Nand:
+ Encoding = bitc::RMW_NAND;
+ break;
+ case AtomicRMWInst::Or:
+ Encoding = bitc::RMW_OR;
+ break;
+ case AtomicRMWInst::Xor:
+ Encoding = bitc::RMW_XOR;
+ break;
+ case AtomicRMWInst::Max:
+ Encoding = bitc::RMW_MAX;
+ break;
+ case AtomicRMWInst::Min:
+ Encoding = bitc::RMW_MIN;
+ break;
+ case AtomicRMWInst::UMax:
+ Encoding = bitc::RMW_UMAX;
+ break;
+ case AtomicRMWInst::UMin:
+ Encoding = bitc::RMW_UMIN;
+ break;
+ case AtomicRMWInst::FAdd:
+ Encoding = bitc::RMW_FADD;
+ break;
+ case AtomicRMWInst::FSub:
+ Encoding = bitc::RMW_FSUB;
+ break;
+ case AtomicRMWInst::FMax:
+ Encoding = bitc::RMW_FMAX;
+ break;
+ case AtomicRMWInst::FMin:
+ Encoding = bitc::RMW_FMIN;
+ break;
case AtomicRMWInst::FMaximum:
- return bitc::RMW_FMAXIMUM;
+ Encoding = bitc::RMW_FMAXIMUM;
+ break;
case AtomicRMWInst::FMinimum:
- return bitc::RMW_FMINIMUM;
+ Encoding = bitc::RMW_FMINIMUM;
+ break;
case AtomicRMWInst::FMaximumNum:
- return bitc::RMW_FMAXIMUMNUM;
+ Encoding = bitc::RMW_FMAXIMUMNUM;
+ break;
case AtomicRMWInst::FMinimumNum:
- return bitc::RMW_FMINIMUMNUM;
+ Encoding = bitc::RMW_FMINIMUMNUM;
+ break;
case AtomicRMWInst::UIncWrap:
- return bitc::RMW_UINC_WRAP;
+ Encoding = bitc::RMW_UINC_WRAP;
+ break;
case AtomicRMWInst::UDecWrap:
- return bitc::RMW_UDEC_WRAP;
+ Encoding = bitc::RMW_UDEC_WRAP;
+ break;
case AtomicRMWInst::USubCond:
- return bitc::RMW_USUB_COND;
+ Encoding = bitc::RMW_USUB_COND;
+ break;
case AtomicRMWInst::USubSat:
- return bitc::RMW_USUB_SAT;
+ Encoding = bitc::RMW_USUB_SAT;
+ break;
}
+
+ if (I.isElementwise())
+ Encoding |= bitc::RMW_ELEMENTWISE_FLAG;
+ return Encoding;
}
static unsigned getEncodedOrdering(AtomicOrdering Ordering) {
@@ -3554,8 +3597,7 @@ void ModuleBitcodeWriter::writeInstruction(const Instruction &I,
Code = bitc::FUNC_CODE_INST_ATOMICRMW;
pushValueAndType(I.getOperand(0), InstID, Vals); // ptrty + ptr
pushValueAndType(I.getOperand(1), InstID, Vals); // valty + val
- Vals.push_back(
- getEncodedRMWOperation(cast<AtomicRMWInst>(I).getOperation()));
+ Vals.push_back(getEncodedRMWOperation(cast<AtomicRMWInst>(I)));
Vals.push_back(cast<AtomicRMWInst>(I).isVolatile());
Vals.push_back(getEncodedOrdering(cast<AtomicRMWInst>(I).getOrdering()));
Vals.push_back(
diff --git a/llvm/lib/IR/AsmWriter.cpp b/llvm/lib/IR/AsmWriter.cpp
index 7bfdb5e9a5b22..826db9212970f 100644
--- a/llvm/lib/IR/AsmWriter.cpp
+++ b/llvm/lib/IR/AsmWriter.cpp
@@ -4483,8 +4483,11 @@ void AssemblyWriter::printInstruction(const Instruction &I) {
Out << ' ' << CI->getPredicate();
// Print out the atomicrmw operation
- if (const auto *RMWI = dyn_cast<AtomicRMWInst>(&I))
+ if (const auto *RMWI = dyn_cast<AtomicRMWInst>(&I)) {
+ if (RMWI->isElementwise())
+ Out << " elementwise";
Out << ' ' << AtomicRMWInst::getOperationName(RMWI->getOperation());
+ }
// Print out the type of the operands...
const Value *Operand = I.getNumOperands() ? I.getOperand(0) : nullptr;
diff --git a/llvm/lib/IR/Instruction.cpp b/llvm/lib/IR/Instruction.cpp
index b002832ec0624..34861c5817d61 100644
--- a/llvm/lib/IR/Instruction.cpp
+++ b/llvm/lib/IR/Instruction.cpp
@@ -969,6 +969,7 @@ bool Instruction::hasSameSpecialState(const Instruction *I2,
cast<AtomicCmpXchgInst>(I2)->getSyncScopeID();
if (const AtomicRMWInst *RMWI = dyn_cast<AtomicRMWInst>(I1))
return RMWI->getOperation() == cast<AtomicRMWInst>(I2)->getOperation() &&
+ RMWI->isElementwise() == cast<AtomicRMWInst>(I2)->isElementwise() &&
RMWI->isVolatile() == cast<AtomicRMWInst>(I2)->isVolatile() &&
(RMWI->getAlign() == cast<AtomicRMWInst>(I2)->getAlign() ||
IgnoreAlignment) &&
diff --git a/llvm/lib/IR/Instructions.cpp b/llvm/lib/IR/Instructions.cpp
index 8a220c48acac8..f940893ab0296 100644
--- a/llvm/lib/IR/Instructions.cpp
+++ b/llvm/lib/IR/Instructions.cpp
@@ -1436,7 +1436,7 @@ AtomicCmpXchgInst::AtomicCmpXchgInst(Value *Ptr, Value *Cmp, Value *NewVal,
void AtomicRMWInst::Init(BinOp Operation, Value *Ptr, Value *Val,
Align Alignment, AtomicOrdering Ordering,
- SyncScope::ID SSID) {
+ SyncScope::ID SSID, bool Elementwise) {
assert(Ordering != AtomicOrdering::NotAtomic &&
"atomicrmw instructions can only be atomic.");
assert(Ordering != AtomicOrdering::Unordered &&
@@ -1446,6 +1446,7 @@ void AtomicRMWInst::Init(BinOp Operation, Value *Ptr, Value *Val,
setOperation(Operation);
setOrdering(Ordering);
setSyncScopeID(SSID);
+ setElementwise(Elementwise);
setAlignment(Alignment);
assert(getOperand(0) && getOperand(1) && "All operands must be non-null!");
@@ -1457,9 +1458,10 @@ void AtomicRMWInst::Init(BinOp Operation, Value *Ptr, Value *Val,
AtomicRMWInst::AtomicRMWInst(BinOp Operation, Value *Ptr, Value *Val,
Align Alignment, AtomicOrdering Ordering,
- SyncScope::ID SSID, InsertPosition InsertBefore)
+ SyncScope::ID SSID, bool Elementwise,
+ InsertPosition InsertBefore)
: Instruction(Val->getType(), AtomicRMW, AllocMarker, InsertBefore) {
- Init(Operation, Ptr, Val, Alignment, Ordering, SSID);
+ Init(Operation, Ptr, Val, Alignment, Ordering, SSID, Elementwise);
}
StringRef AtomicRMWInst::getOperationName(BinOp Op) {
@@ -4397,9 +4399,9 @@ AtomicCmpXchgInst *AtomicCmpXchgInst::cloneImpl() const {
}
AtomicRMWInst *AtomicRMWInst::cloneImpl() const {
- AtomicRMWInst *Result =
- new AtomicRMWInst(getOperation(), getOperand(0), getOperand(1),
- getAlign(), getOrdering(), getSyncScopeID());
+ AtomicRMWInst *Result = new AtomicRMWInst(
+ getOperation(), getOperand(0), getOperand(1), getAlign(), getOrdering(),
+ getSyncScopeID(), isElementwise());
Result->setVolatile(isVolatile());
return Result;
}
diff --git a/llvm/lib/IR/Verifier.cpp b/llvm/lib/IR/Verifier.cpp
index 748e0cc81acfd..2ea113fe665d9 100644
--- a/llvm/lib/IR/Verifier.cpp
+++ b/llvm/lib/IR/Verifier.cpp
@@ -4802,20 +4802,30 @@ void Verifier::visitAtomicRMWInst(AtomicRMWInst &RMWI) {
"atomicrmw instructions cannot be unordered.", &RMWI);
auto Op = RMWI.getOperation();
Type *ElTy = RMWI.getOperand(1)->getType();
+ Type *ScalarTy = ElTy;
+ if (RMWI.isElementwise()) {
+ auto *VecTy = dyn_cast<FixedVectorType>(ElTy);
+ Check(VecTy, "atomicrmw elementwise operand must have fixed vector type!",
+ &RMWI, ElTy);
+ if (VecTy)
+ ScalarTy = VecTy->getElementType();
+ }
+
if (Op == AtomicRMWInst::Xchg) {
- Check(ElTy->isIntegerTy() || ElTy->isFloatingPointTy() ||
- ElTy->isPointerTy(),
+ Check(ScalarTy->isIntegerTy() || ScalarTy->isFloatingPointTy() ||
+ ScalarTy->isPointerTy(),
"atomicrmw " + AtomicRMWInst::getOperationName(Op) +
" operand must have integer or floating point type!",
&RMWI, ElTy);
} else if (AtomicRMWInst::isFPOperation(Op)) {
Check(ElTy->isFPOrFPVectorTy() && !isa<ScalableVectorType>(ElTy),
"atomicrmw " + AtomicRMWInst::getOperationName(Op) +
- " operand must have floating-point or fixed vector of floating-point "
+ " operand must have floating-point or fixed vector of "
+ "floating-point "
"type!",
&RMWI, ElTy);
} else {
- Check(ElTy->isIntegerTy(),
+ Check(ScalarTy->isIntegerTy(),
"atomicrmw " + AtomicRMWInst::getOperationName(Op) +
" operand must have integer type!",
&RMWI, ElTy);
diff --git a/llvm/test/Assembler/atomic.ll b/llvm/test/Assembler/atomic.ll
index 0ed34f0ad98ef..611a717fa9a8d 100644
--- a/llvm/test/Assembler/atomic.ll
+++ b/llvm/test/Assembler/atomic.ll
@@ -151,5 +151,11 @@ define void @fp_vector_atomicrmw(ptr %x, <2 x half> %val) {
; CHECK: %atomic.fminimumnum = atomicrmw fminimumnum ptr %x, <2 x half> %val seq_cst
%atomic.fminimumnum = atomicrmw fminimumnum ptr %x, <2 x half> %val seq_cst
+ ; CHECK: %atomic.elem.fadd = atomicrmw elementwise fadd ptr %x, <2 x half> %val monotonic
+ %atomic.elem.fadd = atomicrmw elementwise fadd ptr %x, <2 x half> %val monotonic
+
+ ; CHECK: %atomic.elem.fadd.vol = atomicrmw volatile elementwise fadd ptr %x, <2 x half> %val seq_cst
+ %atomic.elem.fadd.vol = atomicrmw volatile elementwise fadd ptr %x, <2 x half> %val seq_cst
+
ret void
}
diff --git a/llvm/test/Assembler/invalid-atomicrmw-elementwise.ll b/llvm/test/Assembler/invalid-atomicrmw-elementwise.ll
new file mode 100644
index 0000000000000..a2ac5bed08d94
--- /dev/null
+++ b/llvm/test/Assembler/invalid-atomicrmw-elementwise.ll
@@ -0,0 +1,33 @@
+; RUN: split-file %s %t
+; RUN: not llvm-as -disable-output %t/scalar.ll 2>&1 | FileCheck %t/scalar.ll
+; RUN: not llvm-as -disable-output %t/odd-sized.ll 2>&1 | FileCheck %t/odd-sized.ll
+; RUN: not llvm-as -disable-output %t/add-must-be-integer.ll 2>&1 | FileCheck %t/add-must-be-integer.ll
+; RUN: not llvm-as -disable-output %t/fadd-must-be-fp.ll 2>&1 | FileCheck %t/fadd-must-be-fp.ll
+
+;--- scalar.ll
+; CHECK: atomicrmw elementwise operand must be a fixed vector type
+define i32 @bad_scalar(ptr %p, i32 %v) {
+ %old = atomicrmw elementwise add ptr %p, i32 %v monotonic
+ ret i32 %old
+}
+
+;--- odd-sized.ll
+; CHECK: atomicrmw operand must have a power-of-two byte size
+define <5 x i32> @bad_odd_sized_vector(ptr %p, <5 x i32> %v) {
+ %old = atomicrmw elementwise add ptr %p, <5 x i32> %v monotonic, align 4
+ ret <5 x i32> %old
+}
+
+;--- add-must-be-integer.ll
+; CHECK: atomicrmw add operand must be an integer
+define <4 x float> @bad_add(ptr %p, <4 x float> %v) {
+ %old = atomicrmw elementwise add ptr %p, <4 x float> %v monotonic
+ ret <4 x float> %old
+}
+
+;--- fadd-must-be-fp.ll
+; CHECK: atomicrmw fadd operand must be a floating point type
+define <4 x i32> @bad_fadd(ptr %p, <4 x i32> %v) {
+ %old = atomicrmw elementwise fadd ptr %p, <4 x i32> %v monotonic
+ ret <4 x i32> %old
+}
diff --git a/llvm/test/Bitcode/atomicrmw-elementwise.ll b/llvm/test/Bitcode/atomicrmw-elementwise.ll
new file mode 100644
index 0000000000000..db9c48a80047e
--- /dev/null
+++ b/llvm/test/Bitcode/atomicrmw-elementwise.ll
@@ -0,0 +1,16 @@
+; RUN: llvm-as %s -o - | llvm-dis | FileCheck %s
+; RUN: llvm-as %s -o - | verify-uselistorder
+
+define <4 x i32> @elem_add(ptr %p, <4 x i32> %v) {
+; CHECK-LABEL: @elem_add(
+; CHECK: %old = atomicrmw elementwise add ptr %p, <4 x i32> %v monotonic, align 16
+ %old = atomicrmw elementwise add ptr %p, <4 x i32> %v monotonic
+ ret <4 x i32> %old
+}
+
+define <4 x float> @elem_fadd(ptr %p, <4 x float> %v) {
+; CHECK-LABEL: @elem_fadd(
+; CHECK: %old = atomicrmw elementwise fadd ptr %p, <4 x float> %v seq_cst, align 16
+ %old = atomicrmw elementwise fadd ptr %p, <4 x float> %v seq_cst
+ ret <4 x float> %old
+}
diff --git a/llvm/test/Bitcode/compatibility.ll b/llvm/test/Bitcode/compatibility.ll
index 3c154633740ee..f19b00cc79770 100644
--- a/llvm/test/Bitcode/compatibility.ll
+++ b/llvm/test/Bitcode/compatibility.ll
@@ -1026,6 +1026,16 @@ define void @pointer_atomics(ptr %word) {
ret void
}
+define void @elementwise_atomics(ptr %word, <4 x i32> %ival, <4 x float> %fval) {
+; CHECK: %atomicrmw.add = atomicrmw elementwise add ptr %word, <4 x i32> %ival monotonic, align 16
+ %atomicrmw.add = atomicrmw elementwise add ptr %word, <4 x i32> %ival monotonic, align 16
+
+; CHECK: %atomicrmw.fadd = atomicrmw elementwise fadd ptr %word, <4 x float> %fval seq_cst, align 16
+ %atomicrmw.fadd = atomicrmw elementwise fadd ptr %word, <4 x float> %fval seq_cst, align 16
+
+ ret void
+}
+
;; Fast Math Flags
define void @fastmathflags_unop(float %op1) {
%f.nnan = fneg nnan float %op1
diff --git a/llvm/unittests/Analysis/AliasAnalysisTest.cpp b/llvm/unittests/Analysis/AliasAnalysisTest.cpp
index a28d318ab32c8..f80328b877736 100644
--- a/llvm/unittests/Analysis/AliasAnalysisTest.cpp
+++ b/llvm/unittests/Analysis/AliasAnalysisTest.cpp
@@ -185,7 +185,7 @@ TEST_F(AliasAnalysisTest, getModRefInfo) {
SyncScope::System, BB);
auto *AtomicRMW = new AtomicRMWInst(
AtomicRMWInst::Xchg, Addr, ConstantInt::get(IntType, 1), Alignment,
- AtomicOrdering::Monotonic, SyncScope::System, BB);
+ AtomicOrdering::Monotonic, SyncScope::System, /*Elementwise=*/false, BB);
FunctionType *FooBarTy = FunctionType::get(Type::getVoidTy(C), {}, false);
Function::Create(FooBarTy, Function::ExternalLinkage, "foo", &M);
diff --git a/llvm/unittests/IR/VerifierTest.cpp b/llvm/unittests/IR/VerifierTest.cpp
index e99d2ccd2e548..2d93c57308bcf 100644
--- a/llvm/unittests/IR/VerifierTest.cpp
+++ b/llvm/unittests/IR/VerifierTest.cpp
@@ -374,7 +374,107 @@ TEST(VerifierTest, AtomicRMW) {
Constant *CV = ConstantVector::getSplat(ElementCount::getScalable(2), CF);
new AtomicRMWInst(AtomicRMWInst::FAdd, Ptr, CV, Align(8),
AtomicOrdering::SequentiallyConsistent, SyncScope::System,
- Entry);
+ /*Elementwise=*/false, Entry);
+ ReturnInst::Create(C, Entry);
+
+ std::string Error;
+ raw_string_ostream ErrorOS(Error);
+ EXPECT_TRUE(verifyFunction(*F, &ErrorOS));
+ EXPECT_TRUE(StringRef(Error).starts_with(
+ "atomicrmw fadd operand must have floating-point or "
+ "fixed vector of floating-point type!"))
+ << Error;
+}
+
+TEST(VerifierTest, AtomicRMWElementwiseScalar) {
+ LLVMContext C;
+ Module M("M", C);
+ FunctionType *FTy = FunctionType::get(Type::getVoidTy(C), /*isVarArg=*/false);
+ Function *F = Function::Create(FTy, Function::ExternalLinkage, "foo", M);
+ BasicBlock *Entry = BasicBlock::Create(C, "entry", F);
+ Value *Ptr = PoisonValue::get(PointerType::get(C, 0));
+
+ Type *I32Ty = Type::getInt32Ty(C);
+ Constant *CI = ConstantInt::get(I32Ty, 0);
+
+ new AtomicRMWInst(AtomicRMWInst::Add, Ptr, CI, Align(4),
+ AtomicOrdering::Monotonic, SyncScope::System,
+ /*Elementwise=*/true, Entry);
+ ReturnInst::Create(C, Entry);
+
+ std::string Error;
+ raw_string_ostream ErrorOS(Error);
+ EXPECT_TRUE(verifyFunction(*F, &ErrorOS));
+ EXPECT_TRUE(StringRef(Error).starts_with(
+ "atomicrmw elementwise operand must have fixed vector type!"))
+ << Error;
+}
+
+TEST(VerifierTest, AtomicRMWElementwiseIntOpOnFPVector) {
+ LLVMContext C;
+ Module M("M", C);
+ FunctionType *FTy = FunctionType::get(Type::getVoidTy(C), /*isVarArg=*/false);
+ Function *F = Function::Create(FTy, Function::ExternalLinkage, "foo", M);
+ BasicBlock *Entry = BasicBlock::Create(C, "entry", F);
+ Value *Ptr = PoisonValue::get(PointerType::get(C, 0));
+
+ Type *FPTy = Type::getFloatTy(C);
+ Constant *CV = ConstantVector::getSplat(ElementCount::getFixed(4),
+ ConstantFP::getZero(FPTy));
+
+ new AtomicRMWInst(AtomicRMWInst::Add, Ptr, CV, Align(16),
+ AtomicOrdering::Monotonic, SyncScope::System,
+ /*Elementwise=*/true, Entry);
+ ReturnInst::Create(C, Entry);
+
+ std::string Error;
+ raw_string_ostream ErrorOS(Error);
+ EXPECT_TRUE(verifyFunction(*F, &ErrorOS));
+ EXPECT_TRUE(
+ StringRef(Error).starts_with("atomicrmw add operand must have integer"))
+ << Error;
+}
+
+TEST(VerifierTest, AtomicRMWElementwiseOddSizedVector) {
+ LLVMContext C;
+ Module M("M", C);
+ FunctionType *FTy = FunctionType::get(Type::getVoidTy(C), /*isVarArg=*/false);
+ Function *F = Function::Create(FTy, Function::ExternalLinkage, "foo", M);
+ BasicBlock *Entry = BasicBlock::Create(C, "entry", F);
+ Value *Ptr = PoisonValue::get(PointerType::get(C, 0));
+
+ Type *I32Ty = Type::getInt32Ty(C);
+ Constant *CV = ConstantVector::getSplat(ElementCount::getFixed(5),
+ ConstantInt::get(I32Ty, 0));
+
+ new AtomicRMWInst(AtomicRMWInst::Add, Ptr, CV, Align(4),
+ AtomicOrdering::Monotonic, SyncScope::System,
+ /*Elementwise=*/true, Entry);
+ ReturnInst::Create(C, Entry);
+
+ std::string Error;
+ raw_string_ostream ErrorOS(Error);
+ EXPECT_TRUE(verifyFunction(*F, &ErrorOS));
+ EXPECT_TRUE(StringRef(Error).starts_with(
+ "atomic memory access' operand must have a power-of-two size"))
+ << Error;
+}
+
+TEST(VerifierTest, AtomicRMWElementwiseFPOpOnIntVector) {
+ LLVMContext C;
+ Module M("M", C);
+ FunctionType *FTy = FunctionType::get(Type::getVoidTy(C), /*isVarArg=*/false);
+ Function *F = Function::Create(FTy, Function::ExternalLinkage, "foo", M);
+ BasicBlock *Entry = BasicBlock::Create(C, "entry", F);
+ Value *Ptr = PoisonValue::get(PointerType::get(C, 0));
+
+ Type *I32Ty = Type::getInt32Ty(C);
+ Constant *CV = ConstantVector::getSplat(ElementCount::getFixed(4),
+ ConstantInt::get(I32Ty, 0));
+
+ new AtomicRMWInst(AtomicRMWInst::FAdd, Ptr, CV, Align(16),
+ AtomicOrdering::Monotonic, SyncScope::System,
+ /*Elementwise=*/true, Entry);
ReturnInst::Create(C, Entry);
std::string Error;
More information about the llvm-commits
mailing list